Skip to content
KernelIndex
Search⌘K

submission 108012

macto · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:5b191cad38cefc04e30c9d8d6aa59c0585693b1f5616dd9cd450898fb673c7ed
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))
fp4NVFP4 GEMV - M-tiled Kernel (Non-Persistent, High Parallelism)
fp8__nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
persistent-kernelNVFP4 GEMV - M-tiled Kernel (Non-Persistent, High Parallelism)
shared-memory__shared__ __align__(16) uint8_t sh_a[NUM_BUFFERS][M_TILE][BYTES_PER_TILE];
vector-width = float2float2 f0 = __half22float2(acc0);

Kernel source

submission.py335 lines
"""
NVFP4 GEMV - M-tiled Kernel (Non-Persistent, High Parallelism)

Key idea: Same as submission_1126_rawcuda.py, but each block processes 4 rows.
- Original: grid(M, L), 1 row per block
- This: grid(M/4, L), 4 rows per block with 4 warps

Grid: (M/4, L) = (1792, 1) for M=7168
- Each block: 4 warps, each warp computes 1 row
- B loaded once per block, shared by 4 rows (4× reuse)
- Same high parallelism as original!
"""

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>

// ============================================================================
// Configuration
// ============================================================================
#define BLOCK_SIZE 128             // 4 warps (vs 32 in original)
#define K_TILE 2560
#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;")

// ============================================================================
// FP4/FP8 Conversion
// ============================================================================
__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 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;
}

__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;
}

// ============================================================================
// M-tiled GEMV kernel - 4 warps, each handles one row
// Grid: (M/4, L) - same parallelism as original, with 4x B reuse
// ============================================================================
__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,
    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;
    const int batch_id = blockIdx.y;
    const int tid = threadIdx.x;
    const int warp_id = tid >> 5;   // 0-3: which row this warp handles
    const int lane = tid & 31;      // 0-31: lane within warp

    // 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;  // This warp's row

    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;

    float acc = 0.0f;
    int buf = 0;

    // Load function: B shared, each warp loads its own A row
    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;
        
        // B: All 128 threads cooperate to load (fast!)
        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);
        }
        
        // A: Each warp loads its own row IN PARALLEL
        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);
            }

            // Each warp computes its row
            if (my_row_valid) {
                float tile_acc = 0.0f;
                #pragma unroll 5
                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;
            }
        }
    }

    // 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;
        }
    }

    // Output: each warp writes its own row
    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);
        }
    }
}

// ============================================================================
// 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
) {
    // Grid: (M/4, L) - same parallelism as original grid(M, L)!
    int num_m_blocks = (m + M_TILE - 1) / M_TILE;  // ceil(M/4)
    
    dim3 grid(num_m_blocks, l);  // (1792, 1) for M=7168, L=1
    dim3 block(BLOCK_SIZE);      // 128 threads (4 warps)

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

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_mtiled_v3',
            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 · 335 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 106932.

+ """
+ NVFP4 GEMV - M-tiled Kernel (Non-Persistent, High Parallelism)
+
+ Key idea: Same as submission_1126_rawcuda.py, but each block processes 4 rows.
+ - Original: grid(M, L), 1 row per block
+ - This: grid(M/4, L), 4 rows per block with 4 warps
+
+ Grid: (M/4, L) = (1792, 1) for M=7168
+ - Each block: 4 warps, each warp computes 1 row
+ - B loaded once per block, shared by 4 rows (4× reuse)
+ - Same high parallelism as original!
+ """
+
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
⋯ 4 unchanged lines
#include <cuda_fp4.h>
#include <cuda_fp8.h>
- #define BLOCK_SIZE 32
- #define K_TILE 2560 // Larger tile = fewer iterations
+ // ============================================================================
+ // Configuration
+ // ============================================================================
+ #define BLOCK_SIZE 128 // 4 warps (vs 32 in original)
+ #define K_TILE 2560
#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;")
+
+ // ============================================================================
+ // FP4/FP8 Conversion
+ // ============================================================================
__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
- );
+ static_cast<__nv_fp4x2_storage_t>(byte), __NV_E2M1);
return *reinterpret_cast<__half2*>(&raw);
}
⋯ 2 unchanged lines
__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);
-
+ __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;
}
- __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;
+ __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;
}
// ============================================================================
- // Compute remainder from global memory
+ // M-tiled GEMV kernel - 4 warps, each handles one row
+ // Grid: (M/4, L) - same parallelism as original, with 4x B reuse
// ============================================================================
- __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(
+ void gemv_mtiled_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
+ 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 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_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;
+ // 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;
+ const int batch_id = blockIdx.y;
const int tid = threadIdx.x;
+ const int warp_id = tid >> 5; // 0-3: which row this warp handles
+ const int lane = tid & 31; // 0-31: lane within warp
- if (m >= M) return;
+ // 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; // This warp's row
- // Dimension calculations
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;
- // 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;
+ 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;
- // Tile counts
- const int tile_count = K_bytes / BYTES_PER_TILE;
- const int remainder_start = tile_count * BYTES_PER_TILE;
-
float acc = 0.0f;
+ int buf = 0;
- // 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
+ // Load function: B shared, each warp loads its own A row
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]);
+
+ // B: All 128 threads cooperate to load (fast!)
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;");
+
+ // A: Each warp loads its own row IN PARALLEL
+ 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, 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]);
+ 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_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);
+ 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_sfa_base + i, row_sfa + rem_sf_start + i);
- ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + rem_sf_start + i);
+ ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + remainder_sf_start + i);
}
- asm volatile("cp.async.commit_group;");
+ 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 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;
-
+ // Main K-tile loop with double buffering
if (tile_count > 0) {
- // Load first tile
issue_tile_async(0, 0);
- asm volatile("cp.async.wait_group 0;");
+ ASYNC_WAIT_ALL();
__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);
+ issue_remainder_async(buf ^ 1);
}
- // Compute current tile
- acc += compute_tile(
- sh_a[buf],
- sh_b[buf],
- sh_sfa[buf],
- sh_sfb[buf],
- tid
- );
+ // Each warp computes its row
+ if (my_row_valid) {
+ float tile_acc = 0.0f;
+ #pragma unroll 5
+ 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) {
- asm volatile("cp.async.wait_group 0;");
+ ASYNC_WAIT_ALL();
__syncthreads();
buf ^= 1;
}
}
}
+ // Remainder
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
- );
+ 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;
+ }
}
- #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
- // ====================================================================
+ // Output: each warp writes its own row
+ if (my_row_valid) {
+ float warp_sum = warp_reduce_sum(acc);
if (lane == 0) {
- size_t c_idx = (size_t)m + (size_t)l * M;
- c[c_idx] = __float2half(block_sum);
+ c[(size_t)my_m + (size_t)batch_id * M] = __float2half(warp_sum);
}
}
}
⋯ 2 unchanged lines
// Host wrapper
// ============================================================================
void run_nvfp4_gemv(
- torch::Tensor C,
- torch::Tensor A,
- torch::Tensor B,
- torch::Tensor SFA,
- torch::Tensor SFB,
+ 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);
+ // Grid: (M/4, L) - same parallelism as original grid(M, L)!
+ int num_m_blocks = (m + M_TILE - 1) / M_TILE; // ceil(M/4)
+
+ dim3 grid(num_m_blocks, l); // (1792, 1) for M=7168, L=1
+ dim3 block(BLOCK_SIZE); // 128 threads (4 warps)
- gemv_nvfp4_kernel<<<grid, block>>>(
+ 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()),
⋯ 6 unchanged lines
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
- );
+ 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
⋯ 2 unchanged lines
global _cuda_module
if _cuda_module is None:
_cuda_module = load_inline(
- name='nvfp4_gemv_async_v2',
+ name='nvfp4_gemv_mtiled_v3',
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
-
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
- )
-
+ 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 · 620 diff lines total

Best evidence level for this revision: reported

JSON