Skip to content
KernelIndex
Search⌘K

submission 106932

macto · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:302a2b750e37154d2d5ad64a3ad598535c9886c39f438a968ef267ad7da347bb
license declaredunknown
license concludedunknown
authorsmacto
imported2026-08-15

Techniques

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

async-copyasm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))
fp8__nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
shared-memory__shared__ __align__(16) uint8_t sh_a[NUM_BUFFERS][BYTES_PER_TILE];
vector-width = float2float2 f0 = __half22float2(acc_h2_0);

Kernel source

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

CUDA_SRC = r'''
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>

#define BLOCK_SIZE 32
#define K_TILE 2560          // Larger tile = fewer iterations
#define SCALES_PER_TILE (K_TILE / 16)  // 160
#define BYTES_PER_TILE (K_TILE / 2)    // 1280
#define NUM_BUFFERS 2

__device__ __forceinline__ __half2 decode_fp4x2(uint8_t byte) {
    __half2_raw raw = __nv_cvt_fp4x2_to_halfraw2(
        static_cast<__nv_fp4x2_storage_t>(byte),
        __NV_E2M1
    );
    return *reinterpret_cast<__half2*>(&raw);
}

__device__ __forceinline__ float decode_fp8(int8_t byte) {
    __nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
    __half_raw raw = __nv_cvt_fp8_to_halfraw(storage, __NV_E4M3);
    return __half2float(__ushort_as_half(raw.x));
}
__device__ __forceinline__ __half2 dot_scaled_4bytes(
    uint32_t a4,
    uint32_t b4,
    __half2 scale_h2
) {
    uint32_t b_byte0, b_byte1, b_byte2, b_byte3;
    uint32_t a_byte0, a_byte1, a_byte2, a_byte3;

    asm("bfe.u32 %0, %1, 0, 8;"  : "=r"(b_byte0) : "r"(b4));
    asm("bfe.u32 %0, %1, 8, 8;"  : "=r"(b_byte1) : "r"(b4));
    asm("bfe.u32 %0, %1, 16, 8;" : "=r"(b_byte2) : "r"(b4));
    asm("bfe.u32 %0, %1, 24, 8;" : "=r"(b_byte3) : "r"(b4));

    asm("bfe.u32 %0, %1, 0, 8;"  : "=r"(a_byte0) : "r"(a4));
    asm("bfe.u32 %0, %1, 8, 8;"  : "=r"(a_byte1) : "r"(a4));
    asm("bfe.u32 %0, %1, 16, 8;" : "=r"(a_byte2) : "r"(a4));
    asm("bfe.u32 %0, %1, 24, 8;" : "=r"(a_byte3) : "r"(a4));

    __half2 acc = __hmul2(decode_fp4x2(a_byte0), __hmul2(decode_fp4x2(b_byte0), scale_h2));
    acc = __hfma2(decode_fp4x2(a_byte1), __hmul2(decode_fp4x2(b_byte1), scale_h2), acc);
    acc = __hfma2(decode_fp4x2(a_byte2), __hmul2(decode_fp4x2(b_byte2), scale_h2), acc);
    acc = __hfma2(decode_fp4x2(a_byte3), __hmul2(decode_fp4x2(b_byte3), scale_h2), acc);

    return acc;
}

__device__ __forceinline__ float compute_tile(
    const uint8_t* sh_a,
    const uint8_t* sh_b,
    const uint8_t* sh_sfa,
    const uint8_t* sh_sfb,
    int tid
) {
    float acc = 0.0f;

    #pragma unroll 8
    for (int sf = tid; sf < SCALES_PER_TILE; sf += BLOCK_SIZE) {
        float scale = decode_fp8(static_cast<int8_t>(sh_sfa[sf])) *
                      decode_fp8(static_cast<int8_t>(sh_sfb[sf]));
        __half2 scale_h2 = __half2half2(__float2half(scale));

        int byte_base = sf << 3;  // sf * 8

        uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base]);
        uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base]);
        uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base + 4]);
        uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base + 4]);

        __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
        __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);

        float2 f0 = __half22float2(acc_h2_0);
        float2 f1 = __half22float2(acc_h2_1);
        acc += f0.x + f0.y + f1.x + f1.y;
    }

    return acc;
}

// ============================================================================
// Compute remainder from global memory
// ============================================================================
__device__ __forceinline__ float compute_remainder(
    const uint8_t* row_a,
    const uint8_t* batch_b,
    const uint8_t* row_sfa,
    const uint8_t* batch_sfb,
    int remainder_sf_start,
    int K_sf,
    int tid
) {
    float acc = 0.0f;

    #pragma unroll 4
    for (int sf = remainder_sf_start + tid; sf < K_sf; sf += BLOCK_SIZE) {
        float scale = decode_fp8(static_cast<int8_t>(__ldg(&row_sfa[sf]))) *
                      decode_fp8(static_cast<int8_t>(__ldg(&batch_sfb[sf])));
        __half2 scale_h2 = __half2half2(__float2half(scale));

        int byte_base = sf << 3;

        uint32_t a4_0 = __ldg(reinterpret_cast<const uint32_t*>(&row_a[byte_base]));
        uint32_t b4_0 = __ldg(reinterpret_cast<const uint32_t*>(&batch_b[byte_base]));
        uint32_t a4_1 = __ldg(reinterpret_cast<const uint32_t*>(&row_a[byte_base + 4]));
        uint32_t b4_1 = __ldg(reinterpret_cast<const uint32_t*>(&batch_b[byte_base + 4]));

        __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
        __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);

        float2 f0 = __half22float2(acc_h2_0);
        float2 f1 = __half22float2(acc_h2_1);
        acc += f0.x + f0.y + f1.x + f1.y;
    }

    return acc;
}

__device__ __forceinline__ float compute_remainder_smem(
    const uint8_t* sh_a,
    const uint8_t* sh_b,
    const uint8_t* sh_sfa,
    const uint8_t* sh_sfb,
    int remainder_scales,
    int tid
) {
    float acc = 0.0f;

    #pragma unroll 4
    for (int sf = tid; sf < remainder_scales; sf += BLOCK_SIZE) {
        float scale = decode_fp8(static_cast<int8_t>(sh_sfa[sf])) *
                      decode_fp8(static_cast<int8_t>(sh_sfb[sf]));
        __half2 scale_h2 = __half2half2(__float2half(scale));

        int byte_base = sf << 3;

        uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base]);
        uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base]);
        uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base + 4]);
        uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base + 4]);

        __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
        __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);

        float2 f0 = __half22float2(acc_h2_0);
        float2 f1 = __half22float2(acc_h2_1);
        acc += f0.x + f0.y + f1.x + f1.y;
    }

    return acc;
}

// ============================================================================
// Main GEMV kernel - one block per (row, batch)
// ============================================================================
__global__ __launch_bounds__(BLOCK_SIZE)
void gemv_nvfp4_kernel(
    const int8_t* __restrict__ a,
    const int8_t* __restrict__ b,
    const int8_t* __restrict__ sfa,
    const int8_t* __restrict__ sfb,
    half* __restrict__ c,
    int M, int K, int L,
    int N_rows
) {
    // Double-buffered shared memory
    __shared__ __align__(16) uint8_t sh_a[NUM_BUFFERS][BYTES_PER_TILE];
    __shared__ __align__(16) uint8_t sh_b[NUM_BUFFERS][BYTES_PER_TILE];
    __shared__ __align__(16) uint8_t sh_sfa[NUM_BUFFERS][SCALES_PER_TILE];
    __shared__ __align__(16) uint8_t sh_sfb[NUM_BUFFERS][SCALES_PER_TILE];
    __shared__ float smem_acc[BLOCK_SIZE / 32];

    const int m = blockIdx.x;
    const int l = blockIdx.y;
    const int tid = threadIdx.x;

    if (m >= M) return;

    // Dimension calculations
    const int K_bytes = K / 2;
    const int K_sf = K / 16;

    // Batch strides
    const size_t a_batch_stride = (size_t)M * K_bytes;
    const size_t b_batch_stride = (size_t)N_rows * K_bytes;
    const size_t sfa_batch_stride = (size_t)M * K_sf;
    const size_t sfb_batch_stride = (size_t)N_rows * K_sf;

    // Pointers for this row and batch
    const uint8_t* row_a = reinterpret_cast<const uint8_t*>(a) + l * a_batch_stride + m * K_bytes;
    const uint8_t* batch_b = reinterpret_cast<const uint8_t*>(b) + l * b_batch_stride;
    const uint8_t* row_sfa = reinterpret_cast<const uint8_t*>(sfa) + l * sfa_batch_stride + m * K_sf;
    const uint8_t* batch_sfb = reinterpret_cast<const uint8_t*>(sfb) + l * sfb_batch_stride;

    // Tile counts
    const int tile_count = K_bytes / BYTES_PER_TILE;
    const int remainder_start = tile_count * BYTES_PER_TILE;

    float acc = 0.0f;

    // Async copy macros
    #define ASYNC_COPY_16(dst, src) \
        asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))
    #define ASYNC_COPY_4(dst, src) \
        asm volatile("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"(dst), "l"(src))

    // Lambda for issuing tile async copies
    auto issue_tile_async = [&](int b_idx, int tile) {
        const int base_byte = tile * BYTES_PER_TILE;
        const int base_sf = tile * SCALES_PER_TILE;
        const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][0]);
        const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
        const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][0]);
        const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);

        // Copy A and B tiles (BYTES_PER_TILE bytes each)
        for (int i = tid * 16; i < BYTES_PER_TILE; i += BLOCK_SIZE * 16) {
            ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
            ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
        }
        // Copy scale factors (SCALES_PER_TILE bytes each)
        for (int i = tid * 4; i < SCALES_PER_TILE; i += BLOCK_SIZE * 4) {
            ASYNC_COPY_4(sh_sfa_base + i, row_sfa + base_sf + i);
            ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + base_sf + i);
        }
        asm volatile("cp.async.commit_group;");
    };

    auto issue_remainder_async = [&](int b_idx, int rem_sf_start, int total_K_sf) {
        const int base_byte = rem_sf_start << 3;  // rem_sf_start * 8
        const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][0]);
        const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
        const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][0]);
        const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);

        int remainder_bytes = (total_K_sf - rem_sf_start) << 3;  // * 8
        int remainder_scales = total_K_sf - rem_sf_start;

        for (int i = tid * 16; i < remainder_bytes; i += BLOCK_SIZE * 16) {
            ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
            ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
        }
        for (int i = tid * 4; i < remainder_scales; i += BLOCK_SIZE * 4) {
            ASYNC_COPY_4(sh_sfa_base + i, row_sfa + rem_sf_start + i);
            ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + rem_sf_start + i);
        }
        asm volatile("cp.async.commit_group;");
    };

    // Main loop
    int remainder_sf_start = remainder_start / 8;  // Convert bytes to scale factor index
    bool has_remainder = (remainder_sf_start < K_sf);
    int buf = 0;

    if (tile_count > 0) {
        // Load first tile
        issue_tile_async(0, 0);
        asm volatile("cp.async.wait_group 0;");
        __syncthreads();

        for (int tile = 0; tile < tile_count; ++tile) {
            if (tile + 1 < tile_count) {
                // Prefetch next tile
                issue_tile_async(buf ^ 1, tile + 1);
            } else if (has_remainder) {
                // On last tile: prefetch remainder
                issue_remainder_async(buf ^ 1, remainder_sf_start, K_sf);
            }

            // Compute current tile
            acc += compute_tile(
                sh_a[buf],
                sh_b[buf],
                sh_sfa[buf],
                sh_sfb[buf],
                tid
            );

            if (tile + 1 < tile_count || has_remainder) {
                asm volatile("cp.async.wait_group 0;");
                __syncthreads();
                buf ^= 1;
            }
        }
    }

    if (has_remainder) {
        if (tile_count > 0) {
            int remainder_scales = K_sf - remainder_sf_start;
            acc += compute_remainder_smem(
                sh_a[buf],
                sh_b[buf],
                sh_sfa[buf],
                sh_sfb[buf],
                remainder_scales,
                tid
            );
        } else {
            // No tiles - load remainder directly from global memory
            acc += compute_remainder(
                row_a,
                batch_b,
                row_sfa,
                batch_sfb,
                remainder_sf_start,
                K_sf,
                tid
            );
        }
    }

    #undef ASYNC_COPY_16
    #undef ASYNC_COPY_4

    // ========================================================================
    // Warp-level reduction
    // ========================================================================
    float warp_sum = acc;
    warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 16);
    warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 8);
    warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 4);
    warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 2);
    warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 1);

    // ========================================================================
    // Block-level reduction (only 1 warp, so just write)
    // ========================================================================
    const int warp_id = tid >> 5;
    const int lane = tid & 31;
    
    if (lane == 0) {
        smem_acc[warp_id] = warp_sum;
    }
    __syncthreads();

    if (warp_id == 0) {
        float block_sum = (lane < (BLOCK_SIZE >> 5)) ? smem_acc[lane] : 0.0f;
        block_sum += __shfl_down_sync(0xffffffff, block_sum, 16);
        block_sum += __shfl_down_sync(0xffffffff, block_sum, 8);
        block_sum += __shfl_down_sync(0xffffffff, block_sum, 4);
        block_sum += __shfl_down_sync(0xffffffff, block_sum, 2);
        block_sum += __shfl_down_sync(0xffffffff, block_sum, 1);

        // ====================================================================
        // Final output write
        // ====================================================================
        if (lane == 0) {
            size_t c_idx = (size_t)m + (size_t)l * M;
            c[c_idx] = __float2half(block_sum);
        }
    }
}

// ============================================================================
// Host wrapper
// ============================================================================
void run_nvfp4_gemv(
    torch::Tensor C,
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    int m, int k, int l, int n_pad
) {
    dim3 grid(m, l);
    dim3 block(BLOCK_SIZE);

    gemv_nvfp4_kernel<<<grid, block>>>(
        reinterpret_cast<const int8_t*>(A.data_ptr()),
        reinterpret_cast<const int8_t*>(B.data_ptr()),
        reinterpret_cast<const int8_t*>(SFA.data_ptr()),
        reinterpret_cast<const int8_t*>(SFB.data_ptr()),
        reinterpret_cast<half*>(C.data_ptr<at::Half>()),
        m, k, l, n_pad
    );
}
'''

CPP_SRC = r'''
#include <torch/extension.h>

void run_nvfp4_gemv(
    torch::Tensor C,
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    int m, int k, int l, int n_pad
);
'''

_cuda_module = None

def get_cuda_module():
    global _cuda_module
    if _cuda_module is None:
        _cuda_module = load_inline(
            name='nvfp4_gemv_async_v2',
            cpp_sources=CPP_SRC,
            cuda_sources=CUDA_SRC,
            functions=['run_nvfp4_gemv'],
            extra_cuda_cflags=[
                '-O3',
                '--use_fast_math',
                '-std=c++17',
                '-gencode=arch=compute_100a,code=sm_100a'
            ],
            verbose=False,
        )
    return _cuda_module


def custom_kernel(data: input_t) -> output_t:
    a, b, sfa_ref, sfb_ref, _, _, c = data
    
    module = get_cuda_module()
    
    m, k_packed, l = a.shape
    k = k_packed * 2
    n_pad = b.shape[0]
    
    if not sfa_ref.is_cuda:
        sfa_ref = sfa_ref.to(a.device)
    if not sfb_ref.is_cuda:
        sfb_ref = sfb_ref.to(a.device)
    
    a_uint8 = a.view(torch.uint8)
    b_uint8 = b.view(torch.uint8)
    sfa_uint8 = sfa_ref.view(torch.uint8)
    sfb_uint8 = sfb_ref.view(torch.uint8)
    
    c_out = c.squeeze(1)
    
    module.run_nvfp4_gemv(
        c_out, a_uint8, b_uint8, sfa_uint8, sfb_uint8,
        m, k, l, n_pad
    )
    
    return c
scrolls · 448 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 105811.

- """
- Enhanced GEMV based on our working baseline with selective improvements:
- 1. Native FP4/FP8 conversion functions
- 2. Double buffering for B and SFB with async copy
- 3. Keep our thread block structure (32x32 = 1024 threads)
- 4. Keep our per-thread work (16 bytes per iteration)
- 5. SM100a architecture flag for Blackwell
- """
-
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CUDA_SRC = r'''
#include <torch/extension.h>
- #include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
- #include <cstdint>
- #define THREADS_PER_M 32
- #define THREADS_PER_K 32
- #define BLOCK_SIZE (THREADS_PER_M * THREADS_PER_K)
- #define SF_VEC_SIZE 16
- #define K_TILE_SIZE 512 // 32 threads × 16 bytes
- #define SF_TILE_SIZE 64 // K_TILE_SIZE / 8
+ #define BLOCK_SIZE 32
+ #define K_TILE 2560 // Larger tile = fewer iterations
+ #define SCALES_PER_TILE (K_TILE / 16) // 160
+ #define BYTES_PER_TILE (K_TILE / 2) // 1280
+ #define NUM_BUFFERS 2
- // Native FP4 E2M1 to half2 conversion (from reference)
__device__ __forceinline__ __half2 decode_fp4x2(uint8_t byte) {
__half2_raw raw = __nv_cvt_fp4x2_to_halfraw2(
static_cast<__nv_fp4x2_storage_t>(byte),
⋯ 2 unchanged lines
return *reinterpret_cast<__half2*>(&raw);
}
- // Native FP8 E4M3 to float conversion (from reference)
- __device__ __forceinline__ float decode_fp8(uint8_t byte) {
+ __device__ __forceinline__ float decode_fp8(int8_t byte) {
__nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
__half_raw raw = __nv_cvt_fp8_to_halfraw(storage, __NV_E4M3);
return __half2float(__ushort_as_half(raw.x));
}
+ __device__ __forceinline__ __half2 dot_scaled_4bytes(
+ uint32_t a4,
+ uint32_t b4,
+ __half2 scale_h2
+ ) {
+ uint32_t b_byte0, b_byte1, b_byte2, b_byte3;
+ uint32_t a_byte0, a_byte1, a_byte2, a_byte3;
- // Async copy macros
- #define ASYNC_COPY_16(dst, src) \
- asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))
+ asm("bfe.u32 %0, %1, 0, 8;" : "=r"(b_byte0) : "r"(b4));
+ asm("bfe.u32 %0, %1, 8, 8;" : "=r"(b_byte1) : "r"(b4));
+ asm("bfe.u32 %0, %1, 16, 8;" : "=r"(b_byte2) : "r"(b4));
+ asm("bfe.u32 %0, %1, 24, 8;" : "=r"(b_byte3) : "r"(b4));
- #define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")
- #define ASYNC_WAIT_ALL() asm volatile("cp.async.wait_group 0;")
+ asm("bfe.u32 %0, %1, 0, 8;" : "=r"(a_byte0) : "r"(a4));
+ asm("bfe.u32 %0, %1, 8, 8;" : "=r"(a_byte1) : "r"(a4));
+ asm("bfe.u32 %0, %1, 16, 8;" : "=r"(a_byte2) : "r"(a4));
+ asm("bfe.u32 %0, %1, 24, 8;" : "=r"(a_byte3) : "r"(a4));
- // Kernel for l == 1 with double-buffered B
- __global__ __launch_bounds__(BLOCK_SIZE, 2)
- void nvfp4_gemv_kernel_single(
- half* __restrict__ C,
- const uint8_t* __restrict__ A,
- const uint8_t* __restrict__ B,
- const uint8_t* __restrict__ SFA,
- const uint8_t* __restrict__ SFB,
- int m, int k_packed, int k_sf, int n_pad
+ __half2 acc = __hmul2(decode_fp4x2(a_byte0), __hmul2(decode_fp4x2(b_byte0), scale_h2));
+ acc = __hfma2(decode_fp4x2(a_byte1), __hmul2(decode_fp4x2(b_byte1), scale_h2), acc);
+ acc = __hfma2(decode_fp4x2(a_byte2), __hmul2(decode_fp4x2(b_byte2), scale_h2), acc);
+ acc = __hfma2(decode_fp4x2(a_byte3), __hmul2(decode_fp4x2(b_byte3), scale_h2), acc);
+
+ return acc;
+ }
+
+ __device__ __forceinline__ float compute_tile(
+ const uint8_t* sh_a,
+ const uint8_t* sh_b,
+ const uint8_t* sh_sfa,
+ const uint8_t* sh_sfb,
+ int tid
) {
- // Double buffer for B (shared across all rows)
- __shared__ __align__(16) uint8_t smem_B[2][K_TILE_SIZE];
- __shared__ __align__(16) uint8_t smem_SFB[2][SF_TILE_SIZE];
-
- const int tidx = threadIdx.x;
- const int tidy = threadIdx.y;
- const int tid = tidy * THREADS_PER_K + tidx;
- const int row = blockIdx.x * THREADS_PER_M + tidy;
- const bool valid = (row < m);
-
- const int num_tiles = (k_packed + K_TILE_SIZE - 1) / K_TILE_SIZE;
-
- float sum = 0.0f;
-
- // Load first tile
- if (num_tiles > 0) {
- if (tid < 32) {
- int k_off = tid * 16;
- if (k_off < k_packed) {
- uint32_t dst = __cvta_generic_to_shared(&smem_B[0][tid * 16]);
- ASYNC_COPY_16(dst, B + k_off);
- }
- }
- if (tid < 4) {
- int sf_off = tid * 16;
- if (sf_off < k_sf) {
- uint32_t dst = __cvta_generic_to_shared(&smem_SFB[0][tid * 16]);
- ASYNC_COPY_16(dst, SFB + sf_off);
- }
- }
- ASYNC_COMMIT();
+ float acc = 0.0f;
+
+ #pragma unroll 8
+ for (int sf = tid; sf < SCALES_PER_TILE; sf += BLOCK_SIZE) {
+ float scale = decode_fp8(static_cast<int8_t>(sh_sfa[sf])) *
+ decode_fp8(static_cast<int8_t>(sh_sfb[sf]));
+ __half2 scale_h2 = __half2half2(__float2half(scale));
+
+ int byte_base = sf << 3; // sf * 8
+
+ uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base]);
+ uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base]);
+ uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base + 4]);
+ uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base + 4]);
+
+ __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
+ __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);
+
+ float2 f0 = __half22float2(acc_h2_0);
+ float2 f1 = __half22float2(acc_h2_1);
+ acc += f0.x + f0.y + f1.x + f1.y;
}
-
- for (int tile = 0; tile < num_tiles; tile++) {
- int curr = tile & 1;
- int next = (tile + 1) & 1;
- int next_tile = tile + 1;
-
- // Prefetch next tile
- if (next_tile < num_tiles) {
- int k_base_next = next_tile * K_TILE_SIZE;
- if (tid < 32) {
- int k_off = k_base_next + tid * 16;
- if (k_off < k_packed) {
- uint32_t dst = __cvta_generic_to_shared(&smem_B[next][tid * 16]);
- ASYNC_COPY_16(dst, B + k_off);
- }
- }
- if (tid < 4) {
- int sf_off = k_base_next / 8 + tid * 16;
- if (sf_off < k_sf) {
- uint32_t dst = __cvta_generic_to_shared(&smem_SFB[next][tid * 16]);
- ASYNC_COPY_16(dst, SFB + sf_off);
- }
- }
- ASYNC_COMMIT();
- }
-
- // Wait for current tile
- ASYNC_WAIT_ALL();
- __syncthreads();
-
- if (valid) {
- int k_base = tile * K_TILE_SIZE;
- int k_local = tidx * 16;
- int k_global = k_base + k_local;
-
- if (k_global + 16 <= k_packed) {
- // Load A directly (per-row, can't share)
- uint4 a_vec = *reinterpret_cast<const uint4*>(A + row * k_packed + k_global);
- // Load B from shared memory
- uint4 b_vec = *reinterpret_cast<const uint4*>(&smem_B[curr][k_local]);
-
- const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);
- const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b_vec);
-
- int sf_local = k_local / 8;
- int sf_global = k_global / 8;
-
- float sfa0 = decode_fp8(SFA[row * k_sf + sf_global]);
- float sfa1 = decode_fp8(SFA[row * k_sf + sf_global + 1]);
- float sfb0 = decode_fp8(smem_SFB[curr][sf_local]);
- float sfb1 = decode_fp8(smem_SFB[curr][sf_local + 1]);
- float scale0 = sfa0 * sfb0;
- float scale1 = sfa1 * sfb1;
-
- __half2 acc0 = __float2half2_rn(0.0f);
- __half2 acc1 = __float2half2_rn(0.0f);
-
- #pragma unroll
- for (int i = 0; i < 8; i++) {
- acc0 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc0);
- }
- sum += (__half2float(__low2half(acc0)) + __half2float(__high2half(acc0))) * scale0;
-
- #pragma unroll
- for (int i = 8; i < 16; i++) {
- acc1 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc1);
- }
- sum += (__half2float(__low2half(acc1)) + __half2float(__high2half(acc1))) * scale1;
- }
- }
-
- __syncthreads();
+
+ return acc;
+ }
+
+ // ============================================================================
+ // Compute remainder from global memory
+ // ============================================================================
+ __device__ __forceinline__ float compute_remainder(
+ const uint8_t* row_a,
+ const uint8_t* batch_b,
+ const uint8_t* row_sfa,
+ const uint8_t* batch_sfb,
+ int remainder_sf_start,
+ int K_sf,
+ int tid
+ ) {
+ float acc = 0.0f;
+
+ #pragma unroll 4
+ for (int sf = remainder_sf_start + tid; sf < K_sf; sf += BLOCK_SIZE) {
+ float scale = decode_fp8(static_cast<int8_t>(__ldg(&row_sfa[sf]))) *
+ decode_fp8(static_cast<int8_t>(__ldg(&batch_sfb[sf])));
+ __half2 scale_h2 = __half2half2(__float2half(scale));
+
+ int byte_base = sf << 3;
+
+ uint32_t a4_0 = __ldg(reinterpret_cast<const uint32_t*>(&row_a[byte_base]));
+ uint32_t b4_0 = __ldg(reinterpret_cast<const uint32_t*>(&batch_b[byte_base]));
+ uint32_t a4_1 = __ldg(reinterpret_cast<const uint32_t*>(&row_a[byte_base + 4]));
+ uint32_t b4_1 = __ldg(reinterpret_cast<const uint32_t*>(&batch_b[byte_base + 4]));
+
+ __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
+ __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);
+
+ float2 f0 = __half22float2(acc_h2_0);
+ float2 f1 = __half22float2(acc_h2_1);
+ acc += f0.x + f0.y + f1.x + f1.y;
}
-
- if (!valid) return;
-
- // Warp shuffle reduction
- #pragma unroll
- for (int offset = 16; offset > 0; offset /= 2) {
- sum += __shfl_down_sync(0xffffffff, sum, offset);
+
+ return acc;
+ }
+
+ __device__ __forceinline__ float compute_remainder_smem(
+ const uint8_t* sh_a,
+ const uint8_t* sh_b,
+ const uint8_t* sh_sfa,
+ const uint8_t* sh_sfb,
+ int remainder_scales,
+ int tid
+ ) {
+ float acc = 0.0f;
+
+ #pragma unroll 4
+ for (int sf = tid; sf < remainder_scales; sf += BLOCK_SIZE) {
+ float scale = decode_fp8(static_cast<int8_t>(sh_sfa[sf])) *
+ decode_fp8(static_cast<int8_t>(sh_sfb[sf]));
+ __half2 scale_h2 = __half2half2(__float2half(scale));
+
+ int byte_base = sf << 3;
+
+ uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base]);
+ uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base]);
+ uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base + 4]);
+ uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base + 4]);
+
+ __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
+ __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);
+
+ float2 f0 = __half22float2(acc_h2_0);
+ float2 f1 = __half22float2(acc_h2_1);
+ acc += f0.x + f0.y + f1.x + f1.y;
}
-
- if (tidx == 0) {
- C[row] = __float2half(sum);
- }
+
+ return acc;
}
- // Kernel for l > 1 with double-buffered B
- __global__ __launch_bounds__(BLOCK_SIZE, 2)
- void nvfp4_gemv_kernel_batched(
- half* __restrict__ C,
- const uint8_t* __restrict__ A,
- const uint8_t* __restrict__ B,
- const uint8_t* __restrict__ SFA,
- const uint8_t* __restrict__ SFB,
- int m, int k_packed, int k_sf, int l, int n_pad
+ // ============================================================================
+ // Main GEMV kernel - one block per (row, batch)
+ // ============================================================================
+ __global__ __launch_bounds__(BLOCK_SIZE)
+ void gemv_nvfp4_kernel(
+ const int8_t* __restrict__ a,
+ const int8_t* __restrict__ b,
+ const int8_t* __restrict__ sfa,
+ const int8_t* __restrict__ sfb,
+ half* __restrict__ c,
+ int M, int K, int L,
+ int N_rows
) {
- __shared__ __align__(16) uint8_t smem_B[2][K_TILE_SIZE];
- __shared__ __align__(16) uint8_t smem_SFB[2][SF_TILE_SIZE];
-
- const int tidx = threadIdx.x;
- const int tidy = threadIdx.y;
- const int tid = tidy * THREADS_PER_K + tidx;
- const int row = blockIdx.x * THREADS_PER_M + tidy;
- const int batch = blockIdx.z;
- const bool valid = (row < m);
-
- // Batch offsets
- const size_t a_batch = batch * (size_t)(m * k_packed);
- const size_t b_batch = batch * (size_t)(n_pad * k_packed);
- const size_t sfa_batch = batch * (size_t)(m * k_sf);
- const size_t sfb_batch = batch * (size_t)(n_pad * k_sf);
-
- const int num_tiles = (k_packed + K_TILE_SIZE - 1) / K_TILE_SIZE;
-
- float sum = 0.0f;
-
- // Load first tile
- if (num_tiles > 0) {
- if (tid < 32) {
- int k_off = tid * 16;
- if (k_off < k_packed) {
- uint32_t dst = __cvta_generic_to_shared(&smem_B[0][tid * 16]);
- ASYNC_COPY_16(dst, B + b_batch + k_off);
- }
+ // Double-buffered shared memory
+ __shared__ __align__(16) uint8_t sh_a[NUM_BUFFERS][BYTES_PER_TILE];
+ __shared__ __align__(16) uint8_t sh_b[NUM_BUFFERS][BYTES_PER_TILE];
+ __shared__ __align__(16) uint8_t sh_sfa[NUM_BUFFERS][SCALES_PER_TILE];
+ __shared__ __align__(16) uint8_t sh_sfb[NUM_BUFFERS][SCALES_PER_TILE];
+ __shared__ float smem_acc[BLOCK_SIZE / 32];
+
+ const int m = blockIdx.x;
+ const int l = blockIdx.y;
+ const int tid = threadIdx.x;
+
+ if (m >= M) return;
+
+ // Dimension calculations
+ const int K_bytes = K / 2;
+ const int K_sf = K / 16;
+
+ // Batch strides
+ const size_t a_batch_stride = (size_t)M * K_bytes;
+ const size_t b_batch_stride = (size_t)N_rows * K_bytes;
+ const size_t sfa_batch_stride = (size_t)M * K_sf;
+ const size_t sfb_batch_stride = (size_t)N_rows * K_sf;
+
+ // Pointers for this row and batch
+ const uint8_t* row_a = reinterpret_cast<const uint8_t*>(a) + l * a_batch_stride + m * K_bytes;
+ const uint8_t* batch_b = reinterpret_cast<const uint8_t*>(b) + l * b_batch_stride;
+ const uint8_t* row_sfa = reinterpret_cast<const uint8_t*>(sfa) + l * sfa_batch_stride + m * K_sf;
+ const uint8_t* batch_sfb = reinterpret_cast<const uint8_t*>(sfb) + l * sfb_batch_stride;
+
+ // Tile counts
+ const int tile_count = K_bytes / BYTES_PER_TILE;
+ const int remainder_start = tile_count * BYTES_PER_TILE;
+
+ float acc = 0.0f;
+
+ // Async copy macros
+ #define ASYNC_COPY_16(dst, src) \
+ asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))
+ #define ASYNC_COPY_4(dst, src) \
+ asm volatile("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"(dst), "l"(src))
+
+ // Lambda for issuing tile async copies
+ auto issue_tile_async = [&](int b_idx, int tile) {
+ const int base_byte = tile * BYTES_PER_TILE;
+ const int base_sf = tile * SCALES_PER_TILE;
+ const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][0]);
+ const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
+ const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][0]);
+ const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);
+
+ // Copy A and B tiles (BYTES_PER_TILE bytes each)
+ for (int i = tid * 16; i < BYTES_PER_TILE; i += BLOCK_SIZE * 16) {
+ ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
+ ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
}
- if (tid < 4) {
- int sf_off = tid * 16;
- if (sf_off < k_sf) {
- uint32_t dst = __cvta_generic_to_shared(&smem_SFB[0][tid * 16]);
- ASYNC_COPY_16(dst, SFB + sfb_batch + sf_off);
- }
+ // Copy scale factors (SCALES_PER_TILE bytes each)
+ for (int i = tid * 4; i < SCALES_PER_TILE; i += BLOCK_SIZE * 4) {
+ ASYNC_COPY_4(sh_sfa_base + i, row_sfa + base_sf + i);
+ ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + base_sf + i);
}
- ASYNC_COMMIT();
- }
-
- for (int tile = 0; tile < num_tiles; tile++) {
- int curr = tile & 1;
- int next = (tile + 1) & 1;
- int next_tile = tile + 1;
-
- // Prefetch next tile
- if (next_tile < num_tiles) {
- int k_base_next = next_tile * K_TILE_SIZE;
- if (tid < 32) {
- int k_off = k_base_next + tid * 16;
- if (k_off < k_packed) {
- uint32_t dst = __cvta_generic_to_shared(&smem_B[next][tid * 16]);
- ASYNC_COPY_16(dst, B + b_batch + k_off);
- }
- }
- if (tid < 4) {
- int sf_off = k_base_next / 8 + tid * 16;
- if (sf_off < k_sf) {
- uint32_t dst = __cvta_generic_to_shared(&smem_SFB[next][tid * 16]);
- ASYNC_COPY_16(dst, SFB + sfb_batch + sf_off);
- }
- }
- ASYNC_COMMIT();
+ asm volatile("cp.async.commit_group;");
+ };
+
+ auto issue_remainder_async = [&](int b_idx, int rem_sf_start, int total_K_sf) {
+ const int base_byte = rem_sf_start << 3; // rem_sf_start * 8
+ const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][0]);
+ const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
+ const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][0]);
+ const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);
+
+ int remainder_bytes = (total_K_sf - rem_sf_start) << 3; // * 8
+ int remainder_scales = total_K_sf - rem_sf_start;
+
+ for (int i = tid * 16; i < remainder_bytes; i += BLOCK_SIZE * 16) {
+ ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
+ ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
}
-
- ASYNC_WAIT_ALL();
+ for (int i = tid * 4; i < remainder_scales; i += BLOCK_SIZE * 4) {
+ ASYNC_COPY_4(sh_sfa_base + i, row_sfa + rem_sf_start + i);
+ ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + rem_sf_start + i);
+ }
+ asm volatile("cp.async.commit_group;");
+ };
+
+ // Main loop
+ int remainder_sf_start = remainder_start / 8; // Convert bytes to scale factor index
+ bool has_remainder = (remainder_sf_start < K_sf);
+ int buf = 0;
+
+ if (tile_count > 0) {
+ // Load first tile
+ issue_tile_async(0, 0);
+ asm volatile("cp.async.wait_group 0;");
__syncthreads();
-
- if (valid) {
- int k_base = tile * K_TILE_SIZE;
- int k_local = tidx * 16;
- int k_global = k_base + k_local;
-
- if (k_global + 16 <= k_packed) {
- uint4 a_vec = *reinterpret_cast<const uint4*>(A + a_batch + row * k_packed + k_global);
- uint4 b_vec = *reinterpret_cast<const uint4*>(&smem_B[curr][k_local]);
-
- const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);
- const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b_vec);
-
- int sf_local = k_local / 8;
- int sf_global = k_global / 8;
-
- float sfa0 = decode_fp8(SFA[sfa_batch + row * k_sf + sf_global]);
- float sfa1 = decode_fp8(SFA[sfa_batch + row * k_sf + sf_global + 1]);
- float sfb0 = decode_fp8(smem_SFB[curr][sf_local]);
- float sfb1 = decode_fp8(smem_SFB[curr][sf_local + 1]);
- float scale0 = sfa0 * sfb0;
- float scale1 = sfa1 * sfb1;
-
- __half2 acc0 = __float2half2_rn(0.0f);
- __half2 acc1 = __float2half2_rn(0.0f);
-
- #pragma unroll
- for (int i = 0; i < 8; i++) {
- acc0 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc0);
- }
- sum += (__half2float(__low2half(acc0)) + __half2float(__high2half(acc0))) * scale0;
-
- #pragma unroll
- for (int i = 8; i < 16; i++) {
- acc1 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc1);
- }
- sum += (__half2float(__low2half(acc1)) + __half2float(__high2half(acc1))) * scale1;
+
+ for (int tile = 0; tile < tile_count; ++tile) {
+ if (tile + 1 < tile_count) {
+ // Prefetch next tile
+ issue_tile_async(buf ^ 1, tile + 1);
+ } else if (has_remainder) {
+ // On last tile: prefetch remainder
+ issue_remainder_async(buf ^ 1, remainder_sf_start, K_sf);
}
+
+ // Compute current tile
+ acc += compute_tile(
+ sh_a[buf],
+ sh_b[buf],
+ sh_sfa[buf],
+ sh_sfb[buf],
+ tid
+ );
+
+ if (tile + 1 < tile_count || has_remainder) {
+ asm volatile("cp.async.wait_group 0;");
+ __syncthreads();
+ buf ^= 1;
+ }
}
-
- __syncthreads();
}
-
- if (!valid) return;
-
- #pragma unroll
- for (int offset = 16; offset > 0; offset /= 2) {
- sum += __shfl_down_sync(0xffffffff, sum, offset);
+
+ if (has_remainder) {
+ if (tile_count > 0) {
+ int remainder_scales = K_sf - remainder_sf_start;
+ acc += compute_remainder_smem(
+ sh_a[buf],
+ sh_b[buf],
+ sh_sfa[buf],
+ sh_sfb[buf],
+ remainder_scales,
+ tid
+ );
+ } else {
+ // No tiles - load remainder directly from global memory
+ acc += compute_remainder(
+ row_a,
+ batch_b,
+ row_sfa,
+ batch_sfb,
+ remainder_sf_start,
+ K_sf,
+ tid
+ );
+ }
}
+
+ #undef ASYNC_COPY_16
+ #undef ASYNC_COPY_4
+
+ // ========================================================================
+ // Warp-level reduction
+ // ========================================================================
+ float warp_sum = acc;
+ warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 16);
+ warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 8);
+ warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 4);
+ warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 2);
+ warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 1);
+
+ // ========================================================================
+ // Block-level reduction (only 1 warp, so just write)
+ // ========================================================================
+ const int warp_id = tid >> 5;
+ const int lane = tid & 31;
- if (tidx == 0) {
- C[batch * m + row] = __float2half(sum);
+ if (lane == 0) {
+ smem_acc[warp_id] = warp_sum;
}
+ __syncthreads();
+
+ if (warp_id == 0) {
+ float block_sum = (lane < (BLOCK_SIZE >> 5)) ? smem_acc[lane] : 0.0f;
+ block_sum += __shfl_down_sync(0xffffffff, block_sum, 16);
+ block_sum += __shfl_down_sync(0xffffffff, block_sum, 8);
+ block_sum += __shfl_down_sync(0xffffffff, block_sum, 4);
+ block_sum += __shfl_down_sync(0xffffffff, block_sum, 2);
+ block_sum += __shfl_down_sync(0xffffffff, block_sum, 1);
+
+ // ====================================================================
+ // Final output write
+ // ====================================================================
+ if (lane == 0) {
+ size_t c_idx = (size_t)m + (size_t)l * M;
+ c[c_idx] = __float2half(block_sum);
+ }
+ }
}
+ // ============================================================================
+ // Host wrapper
+ // ============================================================================
void run_nvfp4_gemv(
torch::Tensor C,
torch::Tensor A,
⋯ 2 unchanged lines
torch::Tensor SFB,
int m, int k, int l, int n_pad
) {
- int k_packed = k / 2;
- int k_sf = k / SF_VEC_SIZE;
-
- half* c_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());
- const uint8_t* a_ptr = A.data_ptr<uint8_t>();
- const uint8_t* b_ptr = B.data_ptr<uint8_t>();
- const uint8_t* sfa_ptr = SFA.data_ptr<uint8_t>();
- const uint8_t* sfb_ptr = SFB.data_ptr<uint8_t>();
-
- int blocks_m = (m + THREADS_PER_M - 1) / THREADS_PER_M;
- dim3 block(THREADS_PER_K, THREADS_PER_M, 1);
+ dim3 grid(m, l);
+ dim3 block(BLOCK_SIZE);
- if (l == 1) {
- dim3 grid(blocks_m, 1, 1);
- nvfp4_gemv_kernel_single<<<grid, block>>>(
- c_ptr, a_ptr, b_ptr, sfa_ptr, sfb_ptr,
- m, k_packed, k_sf, n_pad
- );
- } else {
- dim3 grid(blocks_m, 1, l);
- nvfp4_gemv_kernel_batched<<<grid, block>>>(
- c_ptr, a_ptr, b_ptr, sfa_ptr, sfb_ptr,
- m, k_packed, k_sf, l, n_pad
- );
- }
+ gemv_nvfp4_kernel<<<grid, block>>>(
+ reinterpret_cast<const int8_t*>(A.data_ptr()),
+ reinterpret_cast<const int8_t*>(B.data_ptr()),
+ reinterpret_cast<const int8_t*>(SFA.data_ptr()),
+ reinterpret_cast<const int8_t*>(SFB.data_ptr()),
+ reinterpret_cast<half*>(C.data_ptr<at::Half>()),
+ m, k, l, n_pad
+ );
}
-
- #undef ASYNC_COPY_16
- #undef ASYNC_COMMIT
- #undef ASYNC_WAIT_ALL
'''
CPP_SRC = r'''
⋯ 15 unchanged lines
global _cuda_module
if _cuda_module is None:
_cuda_module = load_inline(
- name='nvfp4_gemv_enhanced_v1',
+ name='nvfp4_gemv_async_v2',
cpp_sources=CPP_SRC,
cuda_sources=CUDA_SRC,
functions=['run_nvfp4_gemv'],
extra_cuda_cflags=[
- '-O3',
+ '-O3',
'--use_fast_math',
'-std=c++17',
'-gencode=arch=compute_100a,code=sm_100a'
scrolls · 694 diff lines total

Best evidence level for this revision: reported

JSON