Skip to content
KernelIndex
Search⌘K

submission 104932

yue · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bef06ee649827f10d1cc1c988437ea8af19edd37d4934d4efd7933de54f42655
license declaredunknown
license concludedunknown
authorsyue
imported2026-08-15

Techniques

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

fp4- Vectorized decodes using PTX asm for FP4 x8 and FP8 x4
shared-memory__shared__ __half smem_results[ROWS_PER_BLOCK];
tile-k = 64constexpr int TILE_K = 64; // Each thread handles 64 elements
vector-width = float4float4 A_data[2]; // 2x16 bytes = 32 bytes

Kernel source

submit_v0_ptx3.py283 lines
import torch
import sys
from torch.utils.cpp_extension import load_inline
from typing import Tuple
from task import input_t, output_t

gemv_cuda_src = """
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cuda_pipeline.h>
#include <cstdint>

// Vectorized decode for FP4 x8 - FIXED VERSION
// cvt.rn.f16x2.e2m1x2 takes .b8 input and produces .b32 output (f16x2)
__device__ __forceinline__ void decode_fp4x8_e2m1_half4(const uint32_t packed, __half2& h0, __half2& h1, __half2& h2, __half2& h3)
{
    uint32_t out0, out1, out2, out3;
    
    // PTX allows wider operands for .b8 instruction types
    // Input is uint32 containing 4 bytes, output is 4x .b32 registers
    asm volatile (
        "{"
        "  .reg .b8 %%b0, %%b1, %%b2, %%b3;\\n"
        "  mov.b32 {%%b0, %%b1, %%b2, %%b3}, %4;\\n"
        "  cvt.rn.f16x2.e2m1x2 %0, %%b0;\\n"
        "  cvt.rn.f16x2.e2m1x2 %1, %%b1;\\n"
        "  cvt.rn.f16x2.e2m1x2 %2, %%b2;\\n"
        "  cvt.rn.f16x2.e2m1x2 %3, %%b3;\\n"
        "}"
        : "=r"(out0), "=r"(out1), "=r"(out2), "=r"(out3)
        : "r"(packed)
    );
    
    h0 = *reinterpret_cast<const __half2*>(&out0);
    h1 = *reinterpret_cast<const __half2*>(&out1);
    h2 = *reinterpret_cast<const __half2*>(&out2);
    h3 = *reinterpret_cast<const __half2*>(&out3);
}

// Vectorized decode for FP8 x4 - FIXED VERSION
// cvt.rn.f16x2.e4m3x2 takes .b16 input and produces .b32 output (f16x2)
__device__ __forceinline__ void decode_fp8x4_e4m3fn_half4(const uint32_t packed, __half& h0, __half& h1, __half& h2, __half& h3)
{
    uint32_t out_low, out_high;
    
    // Input is uint32 containing 2x16-bit values
    asm volatile (
        "{"
        "  .reg .b16 %%low, %%high;\\n"
        "  mov.b32 {%%low, %%high}, %2;\\n"
        "  cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
        "  cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
        "}"
        : "=r"(out_low), "=r"(out_high)
        : "r"(packed)
    );
    
    __half2 h_low = *reinterpret_cast<const __half2*>(&out_low);
    __half2 h_high = *reinterpret_cast<const __half2*>(&out_high);
    h0 = h_low.x;
    h1 = h_low.y;
    h2 = h_high.x;
    h3 = h_high.y;
}

extern "C" __global__
void block_scaled_gemv_fp4_fp8_fp16_vectorized(
    const uint8_t* __restrict__ A, // [l, m, k//2]
    const uint8_t* __restrict__ B, // [l, 128, k//2]
    const uint8_t* __restrict__ SFA, // [l, m, k//16]
    const uint8_t* __restrict__ SFB, // [l, 128, k//16]
    __half* __restrict__ C, // [l, m, 1]
    int M, int K, int L)
{
    constexpr int ROWS_PER_BLOCK = 32;
    constexpr int THREADS_PER_ROW = 16; // 16 threads collaborate on each row
    constexpr int TILE_K = 64; // Each thread handles 64 elements
    constexpr int SF_VEC_SIZE = 16;

    const int tidx = threadIdx.x; // 0-15: which K-segment for this row
    const int tidy = threadIdx.y; // 0-31: which row
    const int block_row = blockIdx.x * ROWS_PER_BLOCK;
    const int batch_idx = blockIdx.z;
    const int global_row = block_row + tidy;

    // Shared memory for coalesced writes
    __shared__ __half smem_results[ROWS_PER_BLOCK];

    if (global_row >= M) return;

    float local_sum = 0.0f;
    const int num_k_tiles = K / TILE_K;

    const uint8_t* B_base = B + batch_idx * (128 * K / 2);
    const uint8_t* SFB_base = SFB + batch_idx * (128 * K / 16);
    const uint8_t* A_row = A + batch_idx * (M * K / 2) + global_row * (K / 2);
    const uint8_t* SFA_row = SFA + batch_idx * (M * K / 16) + global_row * (K / 16);

    // K-PARALLELISM: Each thread handles multiple K-tiles, striding by THREADS_PER_ROW
    #pragma unroll 1
    for (int tile = tidx; tile < num_k_tiles; tile += THREADS_PER_ROW) {
        const int k_offset = tile * TILE_K;

        // Vectorized loads
        float4 A_data[2]; // 2x16 bytes = 32 bytes
        float4 B_data[2];

        const float4* A_ptr = reinterpret_cast<const float4*>(A_row + k_offset / 2);
        const float4* B_ptr = reinterpret_cast<const float4*>(B_base + k_offset / 2);

        A_data[0] = __ldg(A_ptr);
        A_data[1] = __ldg(A_ptr + 1);
        B_data[0] = __ldg(B_ptr);
        B_data[1] = __ldg(B_ptr + 1);

        uint8_t* A_tile_data = reinterpret_cast<uint8_t*>(A_data);
        uint8_t* B_tile_data = reinterpret_cast<uint8_t*>(B_data);

        // Load scale factors (4 bytes per thread)
        uint32_t sfa_vec = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + k_offset / 16));
        uint32_t sfb_vec = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + k_offset / 16));

        // Vectorized decode for all 4 scales at once
        __half sfa_scales[4], sfb_scales[4];
        decode_fp8x4_e4m3fn_half4(sfa_vec, sfa_scales[0], sfa_scales[1], sfa_scales[2], sfa_scales[3]);
        decode_fp8x4_e4m3fn_half4(sfb_vec, sfb_scales[0], sfb_scales[1], sfb_scales[2], sfb_scales[3]);

        // COMPUTE with vectorized decode using PTX asm
        constexpr int NUM_SF_BLOCKS = TILE_K / SF_VEC_SIZE;
        
        #pragma unroll
        for (int sf_block = 0; sf_block < NUM_SF_BLOCKS; sf_block++) {
            const int base_k = sf_block * SF_VEC_SIZE;

            // Use pre-decoded scales
            __half sfa = sfa_scales[sf_block];
            __half sfb = sfb_scales[sf_block];
            __half scale = __hmul(sfa, sfb);
            float block_sum = 0.0f;

            // Process 16 elements (8 bytes packed) using two x8 decodes
            #pragma unroll
            for (int sub = 0; sub < 2; sub++) {
                const int byte_idx = base_k / 2 + sub * 4;

                const uint32_t a_packed = *reinterpret_cast<const uint32_t*>(A_tile_data + byte_idx);
                const uint32_t b_packed = *reinterpret_cast<const uint32_t*>(B_tile_data + byte_idx);

                __half2 a0, a1, a2, a3;
                __half2 b0, b1, b2, b3;
                decode_fp4x8_e2m1_half4(a_packed, a0, a1, a2, a3);
                decode_fp4x8_e2m1_half4(b_packed, b0, b1, b2, b3);

                // Vectorized multiplies
                __half2 p0 = __hmul2(a0, b0);
                __half2 p1 = __hmul2(a1, b1);
                __half2 p2 = __hmul2(a2, b2);
                __half2 p3 = __hmul2(a3, b3);

                // Convert to float and accumulate
                float2 fp0 = __half22float2(p0);
                float2 fp1 = __half22float2(p1);
                float2 fp2 = __half22float2(p2);
                float2 fp3 = __half22float2(p3);

                block_sum += fp0.x + fp0.y;
                block_sum += fp1.x + fp1.y;
                block_sum += fp2.x + fp2.y;
                block_sum += fp3.x + fp3.y;
            }

            local_sum += __half2float(scale) * block_sum;
        }
    }

    // Warp-level reduction across K dimension (16 threads per row)
    constexpr unsigned int FULL_MASK = 0xffff; // Mask for 16 threads
    
    #pragma unroll
    for (int offset = 8; offset > 0; offset >>= 1) {
        local_sum += __shfl_down_sync(FULL_MASK, local_sum, offset, 16);
    }

    // First thread in each row writes to shared memory
    if (tidx == 0) {
        smem_results[tidy] = __float2half(local_sum);
    }
    __syncthreads();

    // Coalesced write
    const int linear_tid = tidx + tidy * THREADS_PER_ROW;
    if (linear_tid < ROWS_PER_BLOCK) {
        const int write_row = block_row + linear_tid;
        if (write_row < M) {
            C[batch_idx * M + write_row] = smem_results[linear_tid];
        }
    }
}

torch::Tensor gemv_fp4_fp8_fp16(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C,
    int M, int K, int L)
{
    TORCH_CHECK(A.is_cuda(), "A must be a CUDA tensor");
    TORCH_CHECK(B.is_cuda(), "B must be a CUDA tensor");
    TORCH_CHECK(SFA.is_cuda(), "SFA must be a CUDA tensor");
    TORCH_CHECK(SFB.is_cuda(), "SFB must be a CUDA tensor");
    TORCH_CHECK(C.is_cuda(), "C must be a CUDA tensor");

    constexpr int ROWS_PER_BLOCK = 32;
    constexpr int THREADS_PER_ROW = 16;

    const dim3 grid((M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK, 1, L);
    const dim3 block(THREADS_PER_ROW, ROWS_PER_BLOCK, 1);

    block_scaled_gemv_fp4_fp8_fp16_vectorized<<<grid, block>>>(
        reinterpret_cast<const uint8_t*>(A.data_ptr()),
        reinterpret_cast<const uint8_t*>(B.data_ptr()),
        reinterpret_cast<const uint8_t*>(SFA.data_ptr()),
        reinterpret_cast<const uint8_t*>(SFB.data_ptr()),
        reinterpret_cast<__half*>(C.data_ptr()),
        M, K, L);

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess)
        throw std::runtime_error(cudaGetErrorString(err));

    return C;
}
"""

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

torch::Tensor gemv_fp4_fp8_fp16(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C,
    int M, int K, int L);
"""

_gemm_module = load_inline(
    name="block_scaled_gemv_vectorized",
    cpp_sources=gemv_cpp_src,
    cuda_sources=gemv_cuda_src,
    functions=["gemv_fp4_fp8_fp16"],
    extra_cuda_cflags=[
        '-O3',
        '--use_fast_math',
        '-std=c++17',
        '--expt-relaxed-constexpr',
        '--maxrregcount=256',
        '-gencode=arch=compute_100a,code=sm_100a',
    ],
    verbose=True,
)

def custom_kernel(data: input_t) -> output_t:
    """
    K-parallel with explicit vectorized loads and warp-level reduction:
    - Keeps original K-parallelism (one thread = multiple tiles)
    - Uses __ldg and float4 for explicit vectorized loads
    - Warp shuffle reduction
    - Vectorized decodes using PTX asm for FP4 x8 and FP8 x4
    """
    a, b, sfa, sfb, _, _, c = data
    m, k_packed, l = a.shape
    k = k_packed * 2

    a_uint8 = a.view(torch.uint8)
    b_uint8 = b.view(torch.uint8)
    sfa_uint8 = sfa.view(torch.uint8)
    sfb_uint8 = sfb.view(torch.uint8)

    _gemm_module.gemv_fp4_fp8_fp16(a_uint8, b_uint8, sfa_uint8, sfb_uint8, c, m, k, l)

    return c
scrolls · 283 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 81004.

- # k: 16384; l: 1; m: 7168; seed: 1111
- # ⏱ 53.4 ± 0.05 µs
- # ⚡ 53.1 µs 🐌 55.3 µs
+ import torch
+ import sys
+ from torch.utils.cpp_extension import load_inline
+ from typing import Tuple
+ from task import input_t, output_t
- # k: 7168; l: 8; m: 4096; seed: 1111
- # ⏱ 56.7 ± 0.07 µs
- # ⚡ 55.1 µs 🐌 58.4 µs
+ gemv_cuda_src = """
+ #include <cuda_fp16.h>
+ #include <cuda_runtime.h>
+ #include <cuda_pipeline.h>
+ #include <cstdint>
- # k: 2048; l: 4; m: 7168; seed: 1111
- # ⏱ 21.1 ± 0.07 µs
- # ⚡ 20.4 µs 🐌 22.7 µs
- # Params:
- # When I use:
- # threads_per_m = 128
- # # Make sure threads_per_m is divisible by 1024
- # threads_per_k = 1024 // threads_per_m
- # mma_tiler_mnk = (threads_per_m, 1, 64)
- # It's optimal for k=7168:
- # k: 16384; l: 1; m: 7168; seed: 1111
- # ⏱ 53.4 ± 0.05 µs
- # ⚡ 53.1 µs 🐌 55.3 µs
+ // Vectorized decode for FP4 x8 - FIXED VERSION
+ // cvt.rn.f16x2.e2m1x2 takes .b8 input and produces .b32 output (f16x2)
+ __device__ __forceinline__ void decode_fp4x8_e2m1_half4(const uint32_t packed, __half2& h0, __half2& h1, __half2& h2, __half2& h3)
+ {
+ uint32_t out0, out1, out2, out3;
+
+ // PTX allows wider operands for .b8 instruction types
+ // Input is uint32 containing 4 bytes, output is 4x .b32 registers
+ asm volatile (
+ "{"
+ " .reg .b8 %%b0, %%b1, %%b2, %%b3;\\n"
+ " mov.b32 {%%b0, %%b1, %%b2, %%b3}, %4;\\n"
+ " cvt.rn.f16x2.e2m1x2 %0, %%b0;\\n"
+ " cvt.rn.f16x2.e2m1x2 %1, %%b1;\\n"
+ " cvt.rn.f16x2.e2m1x2 %2, %%b2;\\n"
+ " cvt.rn.f16x2.e2m1x2 %3, %%b3;\\n"
+ "}"
+ : "=r"(out0), "=r"(out1), "=r"(out2), "=r"(out3)
+ : "r"(packed)
+ );
+
+ h0 = *reinterpret_cast<const __half2*>(&out0);
+ h1 = *reinterpret_cast<const __half2*>(&out1);
+ h2 = *reinterpret_cast<const __half2*>(&out2);
+ h3 = *reinterpret_cast<const __half2*>(&out3);
+ }
- # k: 7168; l: 8; m: 4096; seed: 1111
- # ⏱ 57.3 ± 0.06 µs
- # ⚡ 55.3 µs 🐌 58.4 µs
+ // Vectorized decode for FP8 x4 - FIXED VERSION
+ // cvt.rn.f16x2.e4m3x2 takes .b16 input and produces .b32 output (f16x2)
+ __device__ __forceinline__ void decode_fp8x4_e4m3fn_half4(const uint32_t packed, __half& h0, __half& h1, __half& h2, __half& h3)
+ {
+ uint32_t out_low, out_high;
+
+ // Input is uint32 containing 2x16-bit values
+ asm volatile (
+ "{"
+ " .reg .b16 %%low, %%high;\\n"
+ " mov.b32 {%%low, %%high}, %2;\\n"
+ " cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
+ " cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
+ "}"
+ : "=r"(out_low), "=r"(out_high)
+ : "r"(packed)
+ );
+
+ __half2 h_low = *reinterpret_cast<const __half2*>(&out_low);
+ __half2 h_high = *reinterpret_cast<const __half2*>(&out_high);
+ h0 = h_low.x;
+ h1 = h_low.y;
+ h2 = h_high.x;
+ h3 = h_high.y;
+ }
- # k: 2048; l: 4; m: 7168; seed: 1111
- # ⏱ 22.2 ± 0.05 µs
- # ⚡ 20.4 µs 🐌 22.8 µs
+ extern "C" __global__
+ void block_scaled_gemv_fp4_fp8_fp16_vectorized(
+ const uint8_t* __restrict__ A, // [l, m, k//2]
+ const uint8_t* __restrict__ B, // [l, 128, k//2]
+ const uint8_t* __restrict__ SFA, // [l, m, k//16]
+ const uint8_t* __restrict__ SFB, // [l, 128, k//16]
+ __half* __restrict__ C, // [l, m, 1]
+ int M, int K, int L)
+ {
+ constexpr int ROWS_PER_BLOCK = 32;
+ constexpr int THREADS_PER_ROW = 16; // 16 threads collaborate on each row
+ constexpr int TILE_K = 64; // Each thread handles 64 elements
+ constexpr int SF_VEC_SIZE = 16;
- # When I use:
- # threads_per_m = 64
- # # Make sure threads_per_m is divisible by 1024
- # threads_per_k = 1024 // threads_per_m
- # mma_tiler_mnk = (threads_per_m, 1, 64)
- # It's optimal for k=16384:
- # k: 16384; l: 1; m: 7168; seed: 1111
- # ⏱ 34.1 ± 0.06 µs
- # ⚡ 32.7 µs 🐌 34.9 µs
+ const int tidx = threadIdx.x; // 0-15: which K-segment for this row
+ const int tidy = threadIdx.y; // 0-31: which row
+ const int block_row = blockIdx.x * ROWS_PER_BLOCK;
+ const int batch_idx = blockIdx.z;
+ const int global_row = block_row + tidy;
- # k: 7168; l: 8; m: 4096; seed: 1111
- # ⏱ 57.5 ± 0.06 µs
- # ⚡ 56.4 µs 🐌 58.4 µs
+ // Shared memory for coalesced writes
+ __shared__ __half smem_results[ROWS_PER_BLOCK];
- # k: 2048; l: 4; m: 7168; seed: 1111
- # ⏱ 24.6 ± 0.02 µs
- # ⚡ 23.5 µs 🐌 25.6 µs
+ if (global_row >= M) return;
- # Kernel configuration parameters
- # threads_per_m = 32
- # # Make sure threads_per_m is divisible by 512
- # threads_per_k = 512 // threads_per_m
- # Gives k=16384 best performance.
- # k: 16384; l: 1; m: 7168; seed: 1111
- # ⏱ 33.8 ± 0.07 µs
- # ⚡ 32.7 µs 🐌 35.0 µs
+ float local_sum = 0.0f;
+ const int num_k_tiles = K / TILE_K;
- # k: 7168; l: 8; m: 4096; seed: 1111
- # ⏱ 55.3 ± 0.03 µs
- # ⚡ 54.2 µs 🐌 56.3 µs
+ const uint8_t* B_base = B + batch_idx * (128 * K / 2);
+ const uint8_t* SFB_base = SFB + batch_idx * (128 * K / 16);
+ const uint8_t* A_row = A + batch_idx * (M * K / 2) + global_row * (K / 2);
+ const uint8_t* SFA_row = SFA + batch_idx * (M * K / 16) + global_row * (K / 16);
- # k: 2048; l: 4; m: 7168; seed: 1111
- # ⏱ 20.6 ± 0.02 µs
- # ⚡ 20.4 µs 🐌 22.6 µs
+ // K-PARALLELISM: Each thread handles multiple K-tiles, striding by THREADS_PER_ROW
+ #pragma unroll 1
+ for (int tile = tidx; tile < num_k_tiles; tile += THREADS_PER_ROW) {
+ const int k_offset = tile * TILE_K;
- import torch
- from task import input_t, output_t
+ // Vectorized loads
+ float4 A_data[2]; // 2x16 bytes = 32 bytes
+ float4 B_data[2];
- import cutlass
- import cutlass.cute as cute
- from cutlass.cute.runtime import make_ptr
- import cutlass.utils.blockscaled_layout as blockscaled_utils
- from cutlass.utils import SmemAllocator
+ const float4* A_ptr = reinterpret_cast<const float4*>(A_row + k_offset / 2);
+ const float4* B_ptr = reinterpret_cast<const float4*>(B_base + k_offset / 2);
- # Kernel configuration parameters
- threads_per_m = 32
- # Make sure threads_per_m is divisible by 512
- threads_per_k = 512 // threads_per_m
- mma_tiler_mnk = (threads_per_m, 1, 64)
- ab_dtype = cutlass.Float4E2M1FN # FP4 data type for A and B
- sf_dtype = cutlass.Float8E4M3FN # FP8 data type for scale factors
- c_dtype = cutlass.Float16 # FP16 output type
- sf_vec_size = 16 # Scale factor block size (16 elements share one scale)
+ A_data[0] = __ldg(A_ptr);
+ A_data[1] = __ldg(A_ptr + 1);
+ B_data[0] = __ldg(B_ptr);
+ B_data[1] = __ldg(B_ptr + 1);
- # Helper function for ceiling division
- def ceil_div(a, b):
- return (a + b - 1) // b
+ uint8_t* A_tile_data = reinterpret_cast<uint8_t*>(A_data);
+ uint8_t* B_tile_data = reinterpret_cast<uint8_t*>(B_data);
- # The CuTe reference implementation for NVFP4 block-scaled GEMV
- @cute.kernel
- def kernel(
- mA_mkl: cute.Tensor,
- mB_nkl: cute.Tensor,
- mSFA_mkl: cute.Tensor,
- mSFB_nkl: cute.Tensor,
- mC_mnl: cute.Tensor,
- ):
- # Get CUDA block and thread indices
- bidx, bidy, bidz = cute.arch.block_idx()
- tidx, tidy, _ = cute.arch.thread_idx()
+ // Load scale factors (4 bytes per thread)
+ uint32_t sfa_vec = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + k_offset / 16));
+ uint32_t sfb_vec = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + k_offset / 16));
- # Extract the local tile for input matrix A (shape: [block_M, block_K, rest_M, rest_K, rest_L])
- gA_mkl = cute.local_tile(
- mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
- )
- # Extract the local tile for scale factor tensor for A (same shape as gA_mkl)
- # Here, block_M = (32, 4); block_K = (16, 4)
- gSFA_mkl = cute.local_tile(
- mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)
- )
- # Extract the local tile for input matrix B (shape: [block_N, block_K, rest_N, rest_K, rest_L])
- gB_nkl = cute.local_tile(
- mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
- )
- # Extract the local tile for scale factor tensor for B (same shape as gB_nkl)
- gSFB_nkl = cute.local_tile(
- mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)
- )
- # Extract the local tile for output matrix C (shape: [block_M, block_N, rest_M, rest_N, rest_L])
- gC_mnl = cute.local_tile(
- mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)
- )
+ // Vectorized decode for all 4 scales at once
+ __half sfa_scales[4], sfb_scales[4];
+ decode_fp8x4_e4m3fn_half4(sfa_vec, sfa_scales[0], sfa_scales[1], sfa_scales[2], sfa_scales[3]);
+ decode_fp8x4_e4m3fn_half4(sfb_vec, sfb_scales[0], sfb_scales[1], sfb_scales[2], sfb_scales[3]);
- # Select output element corresponding to this thread and block indices
- tCgC = gC_mnl[tidx, None, bidx, bidy, bidz]
- tCgC = cute.make_tensor(tCgC.iterator, 1)
- res = cute.zeros_like(tCgC, cutlass.Float32)
-
- allocator = SmemAllocator()
- # Allocate a buffer for row sum accumulation in shared memory
- # FIXED: Use stride (1, threads_per_m) to avoid bank conflicts
- # With unit stride in tidx dimension, consecutive threads access consecutive addresses
- # This ensures conflict-free access when threads write their results
- row_sum_buffer = allocator.allocate_tensor(
- element_type=cutlass.Float32,
- layout=cute.make_layout((threads_per_m, threads_per_k), stride=(1, threads_per_m))
- )
-
- k_tile_cnt = gA_mkl.layout[3].shape
- for k_tile in range(tidy, k_tile_cnt, threads_per_k):
- tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]
- tBgB = gB_nkl[0, None, bidy, k_tile, bidz]
- tAgSFA = gSFA_mkl[tidx, (0, None), bidx, k_tile, bidz]
- tBgSFB = gSFB_nkl[0, (0, None), bidy, k_tile, bidz]
-
- tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float16)
- tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float16)
- tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)
- tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)
+ // COMPUTE with vectorized decode using PTX asm
+ constexpr int NUM_SF_BLOCKS = TILE_K / SF_VEC_SIZE;
+ #pragma unroll
+ for (int sf_block = 0; sf_block < NUM_SF_BLOCKS; sf_block++) {
+ const int base_k = sf_block * SF_VEC_SIZE;
- # Load NVFP4 or FP8 values from global memory
- a_val_nvfp4 = tAgA.load()
- b_val_nvfp4 = tBgB.load()
- sfa_val_fp8 = tAgSFA.load()
- sfb_val_fp8 = tBgSFB.load()
+ // Use pre-decoded scales
+ __half sfa = sfa_scales[sf_block];
+ __half sfb = sfb_scales[sf_block];
+ __half scale = __hmul(sfa, sfb);
+ float block_sum = 0.0f;
- # Store the converted values to RMEM CuTe tensors
- tArA.store(a_val_nvfp4.to(cutlass.Float16))
- tBrB.store(b_val_nvfp4.to(cutlass.Float16))
- tArSFA.store(sfa_val_fp8.to(cutlass.Float32))
- tBrSFB.store(sfb_val_fp8.to(cutlass.Float32))
+ // Process 16 elements (8 bytes packed) using two x8 decodes
+ #pragma unroll
+ for (int sub = 0; sub < 2; sub++) {
+ const int byte_idx = base_k / 2 + sub * 4;
- # Iterate over SF vector tiles and compute the scale&matmul accumulation
- for sf_block in cutlass.range_constexpr(mma_tiler_mnk[2] // sf_vec_size):
- tmp = cute.zeros_like(tCgC, cutlass.Float32)
- base = sf_block * sf_vec_size
+ const uint32_t a_packed = *reinterpret_cast<const uint32_t*>(A_tile_data + byte_idx);
+ const uint32_t b_packed = *reinterpret_cast<const uint32_t*>(B_tile_data + byte_idx);
- for offset in cutlass.range_constexpr(sf_vec_size):
- tmp += tArA[base + offset] * tBrB[base + offset]
- res += tArSFA[sf_block] * tBrSFB[sf_block] * tmp
+ __half2 a0, a1, a2, a3;
+ __half2 b0, b1, b2, b3;
+ decode_fp4x8_e2m1_half4(a_packed, a0, a1, a2, a3);
+ decode_fp4x8_e2m1_half4(b_packed, b0, b1, b2, b3);
- row_sum_buffer[(tidx, tidy)] = res[0]
- cute.arch.sync_threads()
-
- if tidy == 0:
- out = cute.zeros_like(tCgC, cutlass.Float32)
- for i in cutlass.range_constexpr(threads_per_k):
- out += row_sum_buffer[(tidx, i)]
+ // Vectorized multiplies
+ __half2 p0 = __hmul2(a0, b0);
+ __half2 p1 = __hmul2(a1, b1);
+ __half2 p2 = __hmul2(a2, b2);
+ __half2 p3 = __hmul2(a3, b3);
- # Store the final float16 result back to global memory
- tCgC.store(out.to(cutlass.Float16))
- return
+ // Convert to float and accumulate
+ float2 fp0 = __half22float2(p0);
+ float2 fp1 = __half22float2(p1);
+ float2 fp2 = __half22float2(p2);
+ float2 fp3 = __half22float2(p3);
- @cute.jit
- def my_kernel(
- a_ptr: cute.Pointer,
- b_ptr: cute.Pointer,
- sfa_ptr: cute.Pointer,
- sfb_ptr: cute.Pointer,
- c_ptr: cute.Pointer,
- problem_size: tuple,
- ):
- """
- Host-side JIT function to prepare tensors and launch GPU kernel.
- """
- m, _, k, l = problem_size
- # Create CuTe Tensor via pointer and problem size.
- a_tensor = cute.make_tensor(
- a_ptr,
- cute.make_layout(
- (m, cute.assume(k, 32), l),
- stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),
- ),
- )
- # We use n=128 to create the torch tensor to do fp4 computation via torch._scaled_mm
- # then copy torch tensor to cute tensor for cute customize kernel computation
- # therefore we need to ensure b_tensor has the right stride with this 128 padded size on n.
- n_padded_128 = 128
- b_tensor = cute.make_tensor(
- b_ptr,
- cute.make_layout(
- (n_padded_128, cute.assume(k, 32), l),
- stride=(cute.assume(k, 32), 1, cute.assume(n_padded_128 * k, 32)),
- ),
- )
- c_tensor = cute.make_tensor(
- c_ptr, cute.make_layout((cute.assume(m, 32), 1, l), stride=(1, 1, m))
- )
- # Convert scale factor tensors to MMA layout
- # The layout matches Tensor Core requirements: (((32, 4), REST_M), ((SF_K, 4), REST_K), (1, REST_L))
- sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)
- sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)
+ block_sum += fp0.x + fp0.y;
+ block_sum += fp1.x + fp1.y;
+ block_sum += fp2.x + fp2.y;
+ block_sum += fp3.x + fp3.y;
+ }
- sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)
- sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)
+ local_sum += __half2float(scale) * block_sum;
+ }
+ }
- # Compute grid dimensions
- # Grid is (M_blocks, 1, L) where:
- # - M_blocks = ceil(M / 128) to cover all output rows
- # - L = batch size
- grid = (
- cute.ceil_div(c_tensor.shape[0], threads_per_m),
- 1,
- c_tensor.shape[2],
- )
+ // Warp-level reduction across K dimension (16 threads per row)
+ constexpr unsigned int FULL_MASK = 0xffff; // Mask for 16 threads
+
+ #pragma unroll
+ for (int offset = 8; offset > 0; offset >>= 1) {
+ local_sum += __shfl_down_sync(FULL_MASK, local_sum, offset, 16);
+ }
- # Launch the CUDA kernel
- kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(
- grid=grid,
- block=[threads_per_m, threads_per_k, 1],
- cluster=(1, 1, 1),
- )
- return
+ // First thread in each row writes to shared memory
+ if (tidx == 0) {
+ smem_results[tidy] = __float2half(local_sum);
+ }
+ __syncthreads();
+ // Coalesced write
+ const int linear_tid = tidx + tidy * THREADS_PER_ROW;
+ if (linear_tid < ROWS_PER_BLOCK) {
+ const int write_row = block_row + linear_tid;
+ if (write_row < M) {
+ C[batch_idx * M + write_row] = smem_results[linear_tid];
+ }
+ }
+ }
- # Global cache for compiled kernel
- _compiled_kernel_cache = None
+ torch::Tensor gemv_fp4_fp8_fp16(
+ torch::Tensor A,
+ torch::Tensor B,
+ torch::Tensor SFA,
+ torch::Tensor SFB,
+ torch::Tensor C,
+ int M, int K, int L)
+ {
+ TORCH_CHECK(A.is_cuda(), "A must be a CUDA tensor");
+ TORCH_CHECK(B.is_cuda(), "B must be a CUDA tensor");
+ TORCH_CHECK(SFA.is_cuda(), "SFA must be a CUDA tensor");
+ TORCH_CHECK(SFB.is_cuda(), "SFB must be a CUDA tensor");
+ TORCH_CHECK(C.is_cuda(), "C must be a CUDA tensor");
+ constexpr int ROWS_PER_BLOCK = 32;
+ constexpr int THREADS_PER_ROW = 16;
- # This function is used to compile the kernel once and cache it and then allow users to
- # run the kernel multiple times to get more accurate timing results.
- def compile_kernel():
- """
- Compile the kernel once and cache it.
- This should be called before any timing measurements.
+ const dim3 grid((M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK, 1, L);
+ const dim3 block(THREADS_PER_ROW, ROWS_PER_BLOCK, 1);
- Returns:
- The compiled kernel function
- """
- global _compiled_kernel_cache
+ block_scaled_gemv_fp4_fp8_fp16_vectorized<<<grid, block>>>(
+ reinterpret_cast<const uint8_t*>(A.data_ptr()),
+ reinterpret_cast<const uint8_t*>(B.data_ptr()),
+ reinterpret_cast<const uint8_t*>(SFA.data_ptr()),
+ reinterpret_cast<const uint8_t*>(SFB.data_ptr()),
+ reinterpret_cast<__half*>(C.data_ptr()),
+ M, K, L);
- if _compiled_kernel_cache is not None:
- return _compiled_kernel_cache
+ cudaError_t err = cudaGetLastError();
+ if (err != cudaSuccess)
+ throw std::runtime_error(cudaGetErrorString(err));
- # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
- a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
- b_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
- c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)
- sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
- sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)
+ return C;
+ }
+ """
- # Compile the kernel
- _compiled_kernel_cache = cute.compile(
- my_kernel, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0)
- )
+ gemv_cpp_src = """
+ #include <torch/extension.h>
- return _compiled_kernel_cache
+ torch::Tensor gemv_fp4_fp8_fp16(
+ torch::Tensor A,
+ torch::Tensor B,
+ torch::Tensor SFA,
+ torch::Tensor SFB,
+ torch::Tensor C,
+ int M, int K, int L);
+ """
+ _gemm_module = load_inline(
+ name="block_scaled_gemv_vectorized",
+ cpp_sources=gemv_cpp_src,
+ cuda_sources=gemv_cuda_src,
+ functions=["gemv_fp4_fp8_fp16"],
+ extra_cuda_cflags=[
+ '-O3',
+ '--use_fast_math',
+ '-std=c++17',
+ '--expt-relaxed-constexpr',
+ '--maxrregcount=256',
+ '-gencode=arch=compute_100a,code=sm_100a',
+ ],
+ verbose=True,
+ )
def custom_kernel(data: input_t) -> output_t:
"""
- Execute the block-scaled GEMV kernel.
-
- This is the main entry point called by the evaluation framework.
- It converts PyTorch tensors to CuTe tensors, launches the kernel,
- and returns the result.
-
- Args:
- data: Tuple of (a, b, sfa_cpu, sfb_cpu, c) PyTorch tensors
- a: [m, k, l] - Input matrix in float4e2m1fn
- b: [1, k, l] - Input vector in float4e2m1fn
- sfa_cpu: [m, k, l] - Scale factors in float8_e4m3fn
- sfb_cpu: [1, k, l] - Scale factors in float8_e4m3fn
- sfa_permuted: [32, 4, rest_m, 4, rest_k, l] - Scale factors in float8_e4m3fn
- sfb_permuted: [32, 4, rest_n, 4, rest_k, l] - Scale factors in float8_e4m3fn
- c: [m, 1, l] - Output vector in float16
-
- Returns:
- Output tensor c with computed GEMV results
+ K-parallel with explicit vectorized loads and warp-level reduction:
+ - Keeps original K-parallelism (one thread = multiple tiles)
+ - Uses __ldg and float4 for explicit vectorized loads
+ - Warp shuffle reduction
+ - Vectorized decodes using PTX asm for FP4 x8 and FP8 x4
"""
- a, b, _, _, sfa_permuted, sfb_permuted, c = data
+ a, b, sfa, sfb, _, _, c = data
+ m, k_packed, l = a.shape
+ k = k_packed * 2
- # Ensure kernel is compiled (will use cached version if available)
- # To avoid the compilation overhead, we compile the kernel once and cache it.
- compiled_func = compile_kernel()
+ a_uint8 = a.view(torch.uint8)
+ b_uint8 = b.view(torch.uint8)
+ sfa_uint8 = sfa.view(torch.uint8)
+ sfb_uint8 = sfb.view(torch.uint8)
- # Get dimensions from MxKxL layout
- m, k, l = a.shape
- # Torch use e2m1_x2 data type, thus k is halved
- k = k * 2
- # GEMV N dimension is always 1
- n = 1
+ _gemm_module.gemv_fp4_fp8_fp16(a_uint8, b_uint8, sfa_uint8, sfb_uint8, c, m, k, l)
- # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer
- a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
- b_ptr = make_ptr(ab_dtype, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
- c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)
-
- sfa_ptr = make_ptr(
- sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
- )
- sfb_ptr = make_ptr(
- sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32
- )
-
- # Execute the compiled kernel
- compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))
-
return c
No newline at end of file
scrolls · 568 diff lines total

Best evidence level for this revision: reported

JSON