Skip to content
KernelIndex
Search⌘K

submission 109911

macto · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:43076367a5cf960e20dffca24824707479a984721ded045e4ad69ebf4afe1827
license declaredunknown
license concludedunknown
authorsmacto
imported2026-08-15

Techniques

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

fp4const uint32_t* a_regs, // 4 x uint32 = 16 bytes = 32 FP4
tile-m = 8constexpr int BLOCK_M = 8;
vector-width = half2__device__ __forceinline__ void fp4x8_to_half2x4(half2* out, uint32_t in) {

Kernel source

submission.py264 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>


__device__ __forceinline__ void ldcs_u32x4(uint32_t* dst, const void* src) {
    asm volatile("ld.global.cs.v4.u32 {%0, %1, %2, %3}, [%4];"
        : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3])
        : "l"(src));
}

// Cached load for B (keep in L1, reused across M rows)
__device__ __forceinline__ void ldca_u32x4(uint32_t* dst, const void* src) {
    asm volatile("ld.global.ca.v4.u32 {%0, %1, %2, %3}, [%4];"
        : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3])
        : "l"(src));
}

__device__ __forceinline__ void ldcs_u16(uint16_t* dst, const void* src) {
    asm volatile("ld.global.cs.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
}

// Cached load for scale factors B
__device__ __forceinline__ void ldca_u16(uint16_t* dst, const void* src) {
    asm volatile("ld.global.ca.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
}


// Convert 4 bytes of FP4x2 to 4 x half2 (8 FP16 values)
__device__ __forceinline__ void fp4x8_to_half2x4(half2* out, uint32_t in) {
    asm volatile(
        "{\n\t"
        ".reg .b8 b0, b1, b2, b3;\n\t"
        "mov.b32 {b0, b1, b2, b3}, %4;\n\t"
        "cvt.rn.f16x2.e2m1x2 %0, b0;\n\t"
        "cvt.rn.f16x2.e2m1x2 %1, b1;\n\t"
        "cvt.rn.f16x2.e2m1x2 %2, b2;\n\t"
        "cvt.rn.f16x2.e2m1x2 %3, b3;\n\t"
        "}"
        : "=r"(reinterpret_cast<int&>(out[0])),
          "=r"(reinterpret_cast<int&>(out[1])),
          "=r"(reinterpret_cast<int&>(out[2])),
          "=r"(reinterpret_cast<int&>(out[3]))
        : "r"(in)
    );
}

// Convert 2 FP8 scale factors to half2
__device__ __forceinline__ half2 fp8x2_to_half2(uint16_t in) {
    int out;
    asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(in));
    return *reinterpret_cast<half2*>(&out);
}

__device__ __forceinline__ float blockscaled_dot_16bytes(
    const uint32_t* a_regs,  // 4 x uint32 = 16 bytes = 32 FP4
    const uint32_t* b_regs,  // 4 x uint32 = 16 bytes = 32 FP4
    half2 combined_scale     // SFA * SFB already fused
) {
    // Unpack to half2
    half2 a_h2[16], b_h2[16];
    fp4x8_to_half2x4(a_h2 + 0, a_regs[0]);
    fp4x8_to_half2x4(a_h2 + 4, a_regs[1]);
    fp4x8_to_half2x4(a_h2 + 8, a_regs[2]);
    fp4x8_to_half2x4(a_h2 + 12, a_regs[3]);
    
    fp4x8_to_half2x4(b_h2 + 0, b_regs[0]);
    fp4x8_to_half2x4(b_h2 + 4, b_regs[1]);
    fp4x8_to_half2x4(b_h2 + 8, b_regs[2]);
    fp4x8_to_half2x4(b_h2 + 12, b_regs[3]);
    
    // Dot product WITHOUT scale (deferred scaling)
    // First 8 half2 use scale.x, next 8 use scale.y
    half2 acc0 = __hmul2(a_h2[0], b_h2[0]);
    half2 acc1 = __hmul2(a_h2[8], b_h2[8]);
    
    #pragma unroll
    for (int i = 1; i < 8; i++) {
        acc0 = __hfma2(a_h2[i], b_h2[i], acc0);
        acc1 = __hfma2(a_h2[8 + i], b_h2[8 + i], acc1);
    }
    
    // Reduce each half2 to half
    half sum0 = __hadd(acc0.x, acc0.y);
    half sum1 = __hadd(acc1.x, acc1.y);
    
    float result = 0.0f;
    asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" : "+f"(result) : "h"(*reinterpret_cast<uint16_t*>(&sum0)), "h"(*reinterpret_cast<uint16_t*>(&combined_scale.x)));
    asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" : "+f"(result) : "h"(*reinterpret_cast<uint16_t*>(&sum1)), "h"(*reinterpret_cast<uint16_t*>(&combined_scale.y)));
    
    return result;
}

template<int BLOCK_M, int THREADS_PER_ROW>
__global__ __launch_bounds__(BLOCK_M * THREADS_PER_ROW)
void gemv_optimized_kernel(
    const uint8_t* __restrict__ a,
    const uint8_t* __restrict__ b,
    const uint8_t* __restrict__ sfa,
    const uint8_t* __restrict__ sfb,
    half* __restrict__ c,
    int M, int K, int L, int N_rows
) {
    constexpr int BLOCK_SIZE = BLOCK_M * THREADS_PER_ROW;
    
    const int m_base = blockIdx.x * BLOCK_M;
    const int batch_id = blockIdx.y;
    const int tid = threadIdx.x;
    
    const int warp_id = tid / THREADS_PER_ROW;
    const int lane = tid % THREADS_PER_ROW;
    
    const int m = m_base + warp_id;
    if (m >= M) return;
    
    const int K_bytes = K / 2;
    const int K_sf = K / 16;
    
    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;
    
    const uint8_t* row_a = a + batch_id * a_batch_stride + m * K_bytes;
    const uint8_t* batch_b = b + batch_id * b_batch_stride;
    const uint8_t* row_sfa = sfa + batch_id * sfa_batch_stride + m * K_sf;
    const uint8_t* batch_sfb = sfb + batch_id * sfb_batch_stride;
    
    float acc = 0.0f;
    
    // Process 2 scale groups per iteration (32 FP4 values = 16 bytes)
    // Each thread handles multiple scale pairs
    const int num_scale_pairs = K_sf / 2;
    
    #pragma unroll 4
    for (int sp = lane; sp < num_scale_pairs; sp += THREADS_PER_ROW) {
        int sf_base = sp * 2;
        int byte_base = sf_base * 8;  // 8 bytes per scale factor
        
        // Load scale factors with cache hints
        uint16_t sfa_raw, sfb_raw;
        ldcs_u16(&sfa_raw, row_sfa + sf_base);
        ldca_u16(&sfb_raw, batch_sfb + sf_base);
        
        // Convert FP8x2 to half2 and fuse scales
        half2 scale_a = fp8x2_to_half2(sfa_raw);
        half2 scale_b = fp8x2_to_half2(sfb_raw);
        half2 combined_scale = __hmul2(scale_a, scale_b);
        
        // Load 16 bytes of A and B with cache hints
        uint32_t a_regs[4], b_regs[4];
        ldcs_u32x4(a_regs, row_a + byte_base);
        ldca_u32x4(b_regs, batch_b + byte_base);
        
        // Compute with deferred scaling
        acc += blockscaled_dot_16bytes(a_regs, b_regs, combined_scale);
    }
    
    // Warp reduction
    #pragma unroll
    for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {
        acc += __shfl_down_sync(0xffffffff, acc, offset);
    }
    
    if (lane == 0) {
        c[(size_t)m + (size_t)batch_id * M] = __float2half(acc);
    }
}

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
) {
    // K-specialized dispatch
    if (k >= 8192) {
        // Large K: 8 rows per block, 32 threads per row
        constexpr int BLOCK_M = 8;
        constexpr int THREADS_PER_ROW = 32;
        int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;
        dim3 grid(num_m_blocks, l);
        dim3 block(BLOCK_M * THREADS_PER_ROW);
        gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<grid, block>>>(
            A.data_ptr<uint8_t>(), B.data_ptr<uint8_t>(),
            SFA.data_ptr<uint8_t>(), SFB.data_ptr<uint8_t>(),
            reinterpret_cast<half*>(C.data_ptr<at::Half>()),
            m, k, l, n_pad);
    } else if (k >= 4096) {
        // Medium-large K: 4 rows per block, 32 threads per row
        constexpr int BLOCK_M = 4;
        constexpr int THREADS_PER_ROW = 32;
        int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;
        dim3 grid(num_m_blocks, l);
        dim3 block(BLOCK_M * THREADS_PER_ROW);
        gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<grid, block>>>(
            A.data_ptr<uint8_t>(), B.data_ptr<uint8_t>(),
            SFA.data_ptr<uint8_t>(), SFB.data_ptr<uint8_t>(),
            reinterpret_cast<half*>(C.data_ptr<at::Half>()),
            m, k, l, n_pad);
    } else {
        // Small K: 8 rows per block, 16 threads per row
        constexpr int BLOCK_M = 8;
        constexpr int THREADS_PER_ROW = 16;
        int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;
        dim3 grid(num_m_blocks, l);
        dim3 block(BLOCK_M * THREADS_PER_ROW);
        gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<grid, block>>>(
            A.data_ptr<uint8_t>(), B.data_ptr<uint8_t>(),
            SFA.data_ptr<uint8_t>(), SFB.data_ptr<uint8_t>(),
            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_optimized_v13',
            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)
    c_out = c.squeeze(1)
    module.run_nvfp4_gemv(c_out, a.view(torch.uint8), b.view(torch.uint8),
                          sfa_ref.view(torch.uint8), sfb_ref.view(torch.uint8),
                          m, k, l, n_pad)
    return c

scrolls · 264 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 109158.

⋯ 7 unchanged lines
#include <cuda_fp4.h>
#include <cuda_fp8.h>
- #define BLOCK_SIZE 128 // 4 warps (vs 32 in original)
- #define K_TILE 2048
- #define SCALES_PER_TILE (K_TILE / 16) // 160
- #define BYTES_PER_TILE (K_TILE / 2) // 1280
- #define M_TILE 4 // 4 rows per block (vs 1 in original)
- #define NUM_BUFFERS 2
- // 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))
- #define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")
- #define ASYNC_WAIT_ALL() asm volatile("cp.async.wait_group 0;")
+ __device__ __forceinline__ void ldcs_u32x4(uint32_t* dst, const void* src) {
+ asm volatile("ld.global.cs.v4.u32 {%0, %1, %2, %3}, [%4];"
+ : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3])
+ : "l"(src));
+ }
- __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);
+ // Cached load for B (keep in L1, reused across M rows)
+ __device__ __forceinline__ void ldca_u32x4(uint32_t* dst, const void* src) {
+ asm volatile("ld.global.ca.v4.u32 {%0, %1, %2, %3}, [%4];"
+ : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3])
+ : "l"(src));
}
- __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__ void ldcs_u16(uint16_t* dst, const void* src) {
+ asm volatile("ld.global.cs.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
}
- __device__ __forceinline__ __half2 dot_scaled_4bytes(uint32_t a4, uint32_t b4, __half2 scale_h2) {
- uint32_t b0, b1, b2, b3, a0, a1, a2, a3;
- asm("bfe.u32 %0, %1, 0, 8;" : "=r"(b0) : "r"(b4));
- asm("bfe.u32 %0, %1, 8, 8;" : "=r"(b1) : "r"(b4));
- asm("bfe.u32 %0, %1, 16, 8;" : "=r"(b2) : "r"(b4));
- asm("bfe.u32 %0, %1, 24, 8;" : "=r"(b3) : "r"(b4));
- asm("bfe.u32 %0, %1, 0, 8;" : "=r"(a0) : "r"(a4));
- asm("bfe.u32 %0, %1, 8, 8;" : "=r"(a1) : "r"(a4));
- asm("bfe.u32 %0, %1, 16, 8;" : "=r"(a2) : "r"(a4));
- asm("bfe.u32 %0, %1, 24, 8;" : "=r"(a3) : "r"(a4));
-
- __half2 acc = __hmul2(decode_fp4x2(a0), __hmul2(decode_fp4x2(b0), scale_h2));
- acc = __hfma2(decode_fp4x2(a1), __hmul2(decode_fp4x2(b1), scale_h2), acc);
- acc = __hfma2(decode_fp4x2(a2), __hmul2(decode_fp4x2(b2), scale_h2), acc);
- acc = __hfma2(decode_fp4x2(a3), __hmul2(decode_fp4x2(b3), scale_h2), acc);
- return acc;
+ // Cached load for scale factors B
+ __device__ __forceinline__ void ldca_u16(uint16_t* dst, const void* src) {
+ asm volatile("ld.global.ca.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
}
- __device__ __forceinline__ float warp_reduce_sum(float val) {
- val += __shfl_down_sync(0xffffffff, val, 16);
- val += __shfl_down_sync(0xffffffff, val, 8);
- val += __shfl_down_sync(0xffffffff, val, 4);
- val += __shfl_down_sync(0xffffffff, val, 2);
- val += __shfl_down_sync(0xffffffff, val, 1);
- return val;
+
+ // Convert 4 bytes of FP4x2 to 4 x half2 (8 FP16 values)
+ __device__ __forceinline__ void fp4x8_to_half2x4(half2* out, uint32_t in) {
+ asm volatile(
+ "{\n\t"
+ ".reg .b8 b0, b1, b2, b3;\n\t"
+ "mov.b32 {b0, b1, b2, b3}, %4;\n\t"
+ "cvt.rn.f16x2.e2m1x2 %0, b0;\n\t"
+ "cvt.rn.f16x2.e2m1x2 %1, b1;\n\t"
+ "cvt.rn.f16x2.e2m1x2 %2, b2;\n\t"
+ "cvt.rn.f16x2.e2m1x2 %3, b3;\n\t"
+ "}"
+ : "=r"(reinterpret_cast<int&>(out[0])),
+ "=r"(reinterpret_cast<int&>(out[1])),
+ "=r"(reinterpret_cast<int&>(out[2])),
+ "=r"(reinterpret_cast<int&>(out[3]))
+ : "r"(in)
+ );
}
- __global__ __launch_bounds__(BLOCK_SIZE)
- void gemv_mtiled_kernel(
- const int8_t* __restrict__ a,
- const int8_t* __restrict__ b,
- const int8_t* __restrict__ sfa,
- const int8_t* __restrict__ sfb,
+ // Convert 2 FP8 scale factors to half2
+ __device__ __forceinline__ half2 fp8x2_to_half2(uint16_t in) {
+ int out;
+ asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(in));
+ return *reinterpret_cast<half2*>(&out);
+ }
+
+ __device__ __forceinline__ float blockscaled_dot_16bytes(
+ const uint32_t* a_regs, // 4 x uint32 = 16 bytes = 32 FP4
+ const uint32_t* b_regs, // 4 x uint32 = 16 bytes = 32 FP4
+ half2 combined_scale // SFA * SFB already fused
+ ) {
+ // Unpack to half2
+ half2 a_h2[16], b_h2[16];
+ fp4x8_to_half2x4(a_h2 + 0, a_regs[0]);
+ fp4x8_to_half2x4(a_h2 + 4, a_regs[1]);
+ fp4x8_to_half2x4(a_h2 + 8, a_regs[2]);
+ fp4x8_to_half2x4(a_h2 + 12, a_regs[3]);
+
+ fp4x8_to_half2x4(b_h2 + 0, b_regs[0]);
+ fp4x8_to_half2x4(b_h2 + 4, b_regs[1]);
+ fp4x8_to_half2x4(b_h2 + 8, b_regs[2]);
+ fp4x8_to_half2x4(b_h2 + 12, b_regs[3]);
+
+ // Dot product WITHOUT scale (deferred scaling)
+ // First 8 half2 use scale.x, next 8 use scale.y
+ half2 acc0 = __hmul2(a_h2[0], b_h2[0]);
+ half2 acc1 = __hmul2(a_h2[8], b_h2[8]);
+
+ #pragma unroll
+ for (int i = 1; i < 8; i++) {
+ acc0 = __hfma2(a_h2[i], b_h2[i], acc0);
+ acc1 = __hfma2(a_h2[8 + i], b_h2[8 + i], acc1);
+ }
+
+ // Reduce each half2 to half
+ half sum0 = __hadd(acc0.x, acc0.y);
+ half sum1 = __hadd(acc1.x, acc1.y);
+
+ float result = 0.0f;
+ asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" : "+f"(result) : "h"(*reinterpret_cast<uint16_t*>(&sum0)), "h"(*reinterpret_cast<uint16_t*>(&combined_scale.x)));
+ asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" : "+f"(result) : "h"(*reinterpret_cast<uint16_t*>(&sum1)), "h"(*reinterpret_cast<uint16_t*>(&combined_scale.y)));
+
+ return result;
+ }
+
+ template<int BLOCK_M, int THREADS_PER_ROW>
+ __global__ __launch_bounds__(BLOCK_M * THREADS_PER_ROW)
+ void gemv_optimized_kernel(
+ const uint8_t* __restrict__ a,
+ const uint8_t* __restrict__ b,
+ const uint8_t* __restrict__ sfa,
+ const uint8_t* __restrict__ sfb,
half* __restrict__ c,
int M, int K, int L, int N_rows
) {
- // Shared memory: 4 A rows + 1 B (shared across all 4 warps!)
- __shared__ __align__(16) uint8_t sh_a[NUM_BUFFERS][M_TILE][BYTES_PER_TILE];
- __shared__ __align__(16) uint8_t sh_sfa[NUM_BUFFERS][M_TILE][SCALES_PER_TILE];
- __shared__ __align__(16) uint8_t sh_b[NUM_BUFFERS][BYTES_PER_TILE];
- __shared__ __align__(16) uint8_t sh_sfb[NUM_BUFFERS][SCALES_PER_TILE];
-
- // Block handles M_TILE consecutive rows: [m_base, m_base+1, m_base+2, m_base+3]
- const int m_base = blockIdx.x * M_TILE;
+ constexpr int BLOCK_SIZE = BLOCK_M * THREADS_PER_ROW;
+
+ const int m_base = blockIdx.x * BLOCK_M;
const int batch_id = blockIdx.y;
const int tid = threadIdx.x;
- const int warp_id = tid >> 5;
- const int lane = tid & 31;
-
- // Bounds
- const int valid_rows = min(M_TILE, M - m_base);
- const bool my_row_valid = (warp_id < valid_rows);
- const int my_m = m_base + warp_id;
-
+
+ const int warp_id = tid / THREADS_PER_ROW;
+ const int lane = tid % THREADS_PER_ROW;
+
+ const int m = m_base + warp_id;
+ if (m >= M) return;
+
const int K_bytes = K / 2;
const int K_sf = K / 16;
- const int tile_count = K_bytes / BYTES_PER_TILE;
- const int remainder_sf_start = (tile_count * BYTES_PER_TILE) / 8;
- const bool has_remainder = (remainder_sf_start < K_sf);
- const int remainder_scales = K_sf - remainder_sf_start;
-
+
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;
-
- const uint8_t* batch_a = reinterpret_cast<const uint8_t*>(a) + batch_id * a_batch_stride;
- const uint8_t* batch_b = reinterpret_cast<const uint8_t*>(b) + batch_id * b_batch_stride;
- const uint8_t* batch_sfa = reinterpret_cast<const uint8_t*>(sfa) + batch_id * sfa_batch_stride;
- const uint8_t* batch_sfb = reinterpret_cast<const uint8_t*>(sfb) + batch_id * sfb_batch_stride;
-
+
+ const uint8_t* row_a = a + batch_id * a_batch_stride + m * K_bytes;
+ const uint8_t* batch_b = b + batch_id * b_batch_stride;
+ const uint8_t* row_sfa = sfa + batch_id * sfa_batch_stride + m * K_sf;
+ const uint8_t* batch_sfb = sfb + batch_id * sfb_batch_stride;
+
float acc = 0.0f;
- int buf = 0;
-
- 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;
+
+ // Process 2 scale groups per iteration (32 FP4 values = 16 bytes)
+ // Each thread handles multiple scale pairs
+ const int num_scale_pairs = K_sf / 2;
+
+ #pragma unroll 4
+ for (int sp = lane; sp < num_scale_pairs; sp += THREADS_PER_ROW) {
+ int sf_base = sp * 2;
+ int byte_base = sf_base * 8; // 8 bytes per scale factor
- const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
- const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);
- for (int i = tid * 16; i < BYTES_PER_TILE; i += BLOCK_SIZE * 16) {
- ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
- }
- for (int i = tid * 4; i < SCALES_PER_TILE; i += BLOCK_SIZE * 4) {
- ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + base_sf + i);
- }
+ // Load scale factors with cache hints
+ uint16_t sfa_raw, sfb_raw;
+ ldcs_u16(&sfa_raw, row_sfa + sf_base);
+ ldca_u16(&sfb_raw, batch_sfb + sf_base);
- if (my_row_valid) {
- const uint8_t* row_a = batch_a + my_m * K_bytes;
- const uint8_t* row_sfa = batch_sfa + my_m * K_sf;
- const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][warp_id][0]);
- const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][warp_id][0]);
- for (int i = lane * 16; i < BYTES_PER_TILE; i += 32 * 16) {
- ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
- }
- for (int i = lane * 4; i < SCALES_PER_TILE; i += 32 * 4) {
- ASYNC_COPY_4(sh_sfa_base + i, row_sfa + base_sf + i);
- }
- }
- ASYNC_COMMIT();
- };
-
- auto issue_remainder_async = [&](int b_idx) {
- const int base_byte = remainder_sf_start << 3;
- const int rem_bytes = remainder_scales << 3;
- const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
- const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);
- for (int i = tid * 16; i < rem_bytes; i += BLOCK_SIZE * 16) {
- 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_sfb_base + i, batch_sfb + remainder_sf_start + i);
- }
- if (my_row_valid) {
- const uint8_t* row_a = batch_a + my_m * K_bytes;
- const uint8_t* row_sfa = batch_sfa + my_m * K_sf;
- const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][warp_id][0]);
- const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][warp_id][0]);
- for (int i = lane * 16; i < rem_bytes; i += 32 * 16) {
- ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
- }
- for (int i = lane * 4; i < remainder_scales; i += 32 * 4) {
- ASYNC_COPY_4(sh_sfa_base + i, row_sfa + remainder_sf_start + i);
- }
- }
- ASYNC_COMMIT();
- };
-
- // Main K-tile loop with double buffering
- if (tile_count > 0) {
- issue_tile_async(0, 0);
- ASYNC_WAIT_ALL();
- __syncthreads();
-
- for (int tile = 0; tile < tile_count; ++tile) {
- if (tile + 1 < tile_count) {
- issue_tile_async(buf ^ 1, tile + 1);
- } else if (has_remainder) {
- issue_remainder_async(buf ^ 1);
- }
-
- if (my_row_valid) {
- float tile_acc = 0.0f;
- #pragma unroll 4
- for (int sf = lane; sf < SCALES_PER_TILE; sf += 32) {
- float scale = decode_fp8(static_cast<int8_t>(sh_sfa[buf][warp_id][sf])) *
- decode_fp8(static_cast<int8_t>(sh_sfb[buf][sf]));
- __half2 scale_h2 = __half2half2(__float2half(scale));
- int byte_base = sf << 3;
- uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[buf][warp_id][byte_base]);
- uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[buf][byte_base]);
- uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[buf][warp_id][byte_base + 4]);
- uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[buf][byte_base + 4]);
- __half2 acc0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
- __half2 acc1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);
- float2 f0 = __half22float2(acc0);
- float2 f1 = __half22float2(acc1);
- tile_acc += f0.x + f0.y + f1.x + f1.y;
- }
- acc += tile_acc;
- }
-
- if (tile + 1 < tile_count || has_remainder) {
- ASYNC_WAIT_ALL();
- __syncthreads();
- buf ^= 1;
- }
- }
+ // Convert FP8x2 to half2 and fuse scales
+ half2 scale_a = fp8x2_to_half2(sfa_raw);
+ half2 scale_b = fp8x2_to_half2(sfb_raw);
+ half2 combined_scale = __hmul2(scale_a, scale_b);
+
+ // Load 16 bytes of A and B with cache hints
+ uint32_t a_regs[4], b_regs[4];
+ ldcs_u32x4(a_regs, row_a + byte_base);
+ ldca_u32x4(b_regs, batch_b + byte_base);
+
+ // Compute with deferred scaling
+ acc += blockscaled_dot_16bytes(a_regs, b_regs, combined_scale);
}
-
- // Remainder
- if (has_remainder) {
- if (tile_count == 0) {
- issue_remainder_async(0);
- ASYNC_WAIT_ALL();
- __syncthreads();
- buf = 0;
- }
- if (my_row_valid) {
- float rem_acc = 0.0f;
- for (int sf = lane; sf < remainder_scales; sf += 32) {
- float scale = decode_fp8(static_cast<int8_t>(sh_sfa[buf][warp_id][sf])) *
- decode_fp8(static_cast<int8_t>(sh_sfb[buf][sf]));
- __half2 scale_h2 = __half2half2(__float2half(scale));
- int byte_base = sf << 3;
- uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[buf][warp_id][byte_base]);
- uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[buf][byte_base]);
- uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[buf][warp_id][byte_base + 4]);
- uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[buf][byte_base + 4]);
- __half2 acc0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
- __half2 acc1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);
- float2 f0 = __half22float2(acc0);
- float2 f1 = __half22float2(acc1);
- rem_acc += f0.x + f0.y + f1.x + f1.y;
- }
- acc += rem_acc;
- }
+
+ // Warp reduction
+ #pragma unroll
+ for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {
+ acc += __shfl_down_sync(0xffffffff, acc, offset);
}
-
- if (my_row_valid) {
- float warp_sum = warp_reduce_sum(acc);
- if (lane == 0) {
- c[(size_t)my_m + (size_t)batch_id * M] = __float2half(warp_sum);
- }
+
+ if (lane == 0) {
+ c[(size_t)m + (size_t)batch_id * M] = __float2half(acc);
}
}
⋯ 2 unchanged lines
torch::Tensor SFA, torch::Tensor SFB,
int m, int k, int l, int n_pad
) {
- int num_m_blocks = (m + M_TILE - 1) / M_TILE;
-
- dim3 grid(num_m_blocks, l);
- dim3 block(BLOCK_SIZE);
-
- gemv_mtiled_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
- );
+ // K-specialized dispatch
+ if (k >= 8192) {
+ // Large K: 8 rows per block, 32 threads per row
+ constexpr int BLOCK_M = 8;
+ constexpr int THREADS_PER_ROW = 32;
+ int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;
+ dim3 grid(num_m_blocks, l);
+ dim3 block(BLOCK_M * THREADS_PER_ROW);
+ gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<grid, block>>>(
+ A.data_ptr<uint8_t>(), B.data_ptr<uint8_t>(),
+ SFA.data_ptr<uint8_t>(), SFB.data_ptr<uint8_t>(),
+ reinterpret_cast<half*>(C.data_ptr<at::Half>()),
+ m, k, l, n_pad);
+ } else if (k >= 4096) {
+ // Medium-large K: 4 rows per block, 32 threads per row
+ constexpr int BLOCK_M = 4;
+ constexpr int THREADS_PER_ROW = 32;
+ int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;
+ dim3 grid(num_m_blocks, l);
+ dim3 block(BLOCK_M * THREADS_PER_ROW);
+ gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<grid, block>>>(
+ A.data_ptr<uint8_t>(), B.data_ptr<uint8_t>(),
+ SFA.data_ptr<uint8_t>(), SFB.data_ptr<uint8_t>(),
+ reinterpret_cast<half*>(C.data_ptr<at::Half>()),
+ m, k, l, n_pad);
+ } else {
+ // Small K: 8 rows per block, 16 threads per row
+ constexpr int BLOCK_M = 8;
+ constexpr int THREADS_PER_ROW = 16;
+ int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;
+ dim3 grid(num_m_blocks, l);
+ dim3 block(BLOCK_M * THREADS_PER_ROW);
+ gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<grid, block>>>(
+ A.data_ptr<uint8_t>(), B.data_ptr<uint8_t>(),
+ SFA.data_ptr<uint8_t>(), SFB.data_ptr<uint8_t>(),
+ reinterpret_cast<half*>(C.data_ptr<at::Half>()),
+ m, k, l, n_pad);
+ }
}
'''
⋯ 9 unchanged lines
global _cuda_module
if _cuda_module is None:
_cuda_module = load_inline(
- name='nvfp4_gemv_mtiled_v3',
+ name='nvfp4_gemv_optimized_v13',
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'],
+ extra_cuda_cflags=[
+ '-O3',
+ '--use_fast_math',
+ '-std=c++17',
+ '-gencode=arch=compute_100a,code=sm_100a',
+ ],
verbose=False,
)
return _cuda_module
⋯ 13 unchanged lines
sfa_ref.view(torch.uint8), sfb_ref.view(torch.uint8),
m, k, l, n_pad)
return c
+
scrolls · 457 diff lines total

Best evidence level for this revision: reported

JSON