Skip to content
KernelIndex
Search⌘K

submission 108659

currybab · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:0ac3f8a6a8cdfc94ebe8540e19daec3852625765dc2375653d369ed2fe232f6b
license declaredunknown
license concludedunknown
authorscurrybab
imported2026-08-15

Techniques

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

fp8__device__ __forceinline__ __half decode_fp8_h(__nv_fp8_storage_t v) {
shared-memory__shared__ float smem_acc[4];
vector-width = uint4uint4 a_vec0, a_vec1, b_vec0, b_vec1;

Kernel source

submission_l1_specialized.py550 lines
import torch
import sys
import os
import random
import re
import numpy as np
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

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

#define K_THREADS 32
#define K_L1_THREADS 128
#define WARP_SIZE 32

// Shared memory B caching parameters
#define M_TILE 8        // Number of M rows per block
#define SMEM_K_THREADS 32  // K-direction threads per M row

__device__ __forceinline__ __half decode_fp8_h(__nv_fp8_storage_t v) {
    return __half(__nv_cvt_fp8_to_halfraw(v, __NV_E4M3));
}
__device__ __forceinline__ __half2 decode_fp4_h2(__nv_fp4x2_storage_t v) {
    return __half2(__nv_cvt_fp4x2_to_halfraw2(v, __NV_E2M1));
}

// =================================================================================================
// Kernel: Specialized L1 Kernel (Optimized for L=1, K=16384)
// =================================================================================================
// Based on v2_merge's gemv_kernel_l1, but with 'l' loop and offsets removed/simplified.
// Hardcoded for L=1.
__global__ void gemv_kernel_l1_specialized(
    const __nv_fp4x2_storage_t* __restrict__ A,
    const __nv_fp4x2_storage_t* __restrict__ B,
    const __nv_fp8_storage_t*  __restrict__ sfa,
    const __nv_fp8_storage_t*  __restrict__ sfb,
    half* __restrict__ C,
    int M, int K // L is implicitly 1
) {
    int m = blockIdx.x;
    // L is always 0, so l_offset is 0.
    
    int tid_k = threadIdx.x;

    // Shared memory for inter-warp reduction
    // K_L1_THREADS / WARP_SIZE = 128 / 32 = 4
    __shared__ float smem_acc[4];

    if (m >= M) return;

    const int K_sf = K / 16;
    const int K_half = K / 2;

    // Simplified offsets (L=0)
    const size_t seg_a = (size_t)m * K_half;
    // seg_b is 0
    const size_t seg_sfa = (size_t)m * K_sf;
    // seg_sfb is 0

    float acc = 0.0f;

    // 4 blocks per thread iteration (Unroll factor 4)
    int k_base = tid_k * 4;
    
    // Pre-load variables
    uint4 a_vec0, a_vec1, b_vec0, b_vec1;
    uint sfa_packed, sfb_packed;

    // Initial Load
    if (k_base < K_sf) {
        sfa_packed = *reinterpret_cast<const uint*>(&sfa[seg_sfa + k_base]);
        sfb_packed = *reinterpret_cast<const uint*>(&sfb[k_base]); // sfb is at offset 0
        a_vec0 = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8]);
        a_vec1 = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8 + 16]); // +16 uint4 offset? No.
        // Wait, previous code: a_vec1 = ... &A[seg_a + k_base * 8 + 16]
        // A is __nv_fp4x2_storage_t* (1 byte).
        // k_base * 8 -> 8 bytes per block.
        // a_vec0 loads 16 bytes (2 blocks: k_base, k_base+1).
        // a_vec1 loads 16 bytes (2 blocks: k_base+2, k_base+3).
        // 2 blocks = 16 bytes.
        // So offset for a_vec1 should be + 16 bytes from a_vec0 address.
        // Pointer arithmetic: A + ... 
        // &A[...] returns address.
        // k_base*8 is byte offset? No, A is typed pointer.
        // sizeof(__nv_fp4x2_storage_t) = 1.
        // So A[...] is byte addressing basically.
        // k_base blocks * 8 bytes/block = k_base*8 bytes.
        // a_vec0 reads bytes [k_base*8 ... k_base*8+15].
        // a_vec1 reads bytes [k_base*8+16 ... k_base*8+31].
        // So +16 is correct.
        
        a_vec1 = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8 + 16]);
        b_vec0 = *reinterpret_cast<const uint4*>(&B[k_base * 8]);
        b_vec1 = *reinterpret_cast<const uint4*>(&B[k_base * 8 + 16]);
    }

    // Main Loop
    // Stride: K_L1_THREADS * 4 = 128 * 4 = 512 blocks
    int stride_blocks = K_L1_THREADS * 4;

    for (; k_base < K_sf; k_base += stride_blocks) {
        uint4 curr_a0 = a_vec0, curr_a1 = a_vec1;
        uint4 curr_b0 = b_vec0, curr_b1 = b_vec1;
        uint curr_sfa = sfa_packed, curr_sfb = sfb_packed;

        // Prefetch next iteration
        int next_k = k_base + stride_blocks;
        if (next_k < K_sf) {
            sfa_packed = *reinterpret_cast<const uint*>(&sfa[seg_sfa + next_k]);
            sfb_packed = *reinterpret_cast<const uint*>(&sfb[next_k]);
            a_vec0 = *reinterpret_cast<const uint4*>(&A[seg_a + next_k * 8]);
            a_vec1 = *reinterpret_cast<const uint4*>(&A[seg_a + next_k * 8 + 16]);
            b_vec0 = *reinterpret_cast<const uint4*>(&B[next_k * 8]);
            b_vec1 = *reinterpret_cast<const uint4*>(&B[next_k * 8 + 16]);
        }

        const __nv_fp8_storage_t* sfa_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&curr_sfa);
        const __nv_fp8_storage_t* sfb_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&curr_sfb);

        // Process first 2 blocks (k_base, k_base+1)
        const __nv_fp4x2_storage_t* a_fp4x2_0 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_a0);
        const __nv_fp4x2_storage_t* b_fp4x2_0 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_b0);

        #pragma unroll
        for (int i = 0; i < 2; i++) {
            if (k_base + i >= K_sf) break; // Safety check, usually not needed if K divisible
            float scale_ab = __half2float(decode_fp8_h(sfa_bytes[i])) * __half2float(decode_fp8_h(sfb_bytes[i]));
            
            __half2 block_acc_h2 = __float2half2_rn(0.0f);
            #pragma unroll
            for (int j = 0; j < 8; j++) {
                __half2 a2 = decode_fp4_h2(a_fp4x2_0[i * 8 + j]);
                __half2 b2 = decode_fp4_h2(b_fp4x2_0[i * 8 + j]);
                block_acc_h2 = __hfma2(a2, b2, block_acc_h2);
            }
            float2 pf = __half22float2(block_acc_h2);
            acc = __fmaf_rn(pf.x + pf.y, scale_ab, acc);
        }

        // Process next 2 blocks (k_base+2, k_base+3)
        const __nv_fp4x2_storage_t* a_fp4x2_1 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_a1);
        const __nv_fp4x2_storage_t* b_fp4x2_1 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_b1);

        #pragma unroll
        for (int i = 0; i < 2; i++) {
            if (k_base + 2 + i >= K_sf) break;
            float scale_ab = __half2float(decode_fp8_h(sfa_bytes[2 + i])) * __half2float(decode_fp8_h(sfb_bytes[2 + i]));
            
            __half2 block_acc_h2 = __float2half2_rn(0.0f);
            #pragma unroll
            for (int j = 0; j < 8; j++) {
                __half2 a2 = decode_fp4_h2(a_fp4x2_1[i * 8 + j]);
                __half2 b2 = decode_fp4_h2(b_fp4x2_1[i * 8 + j]);
                block_acc_h2 = __hfma2(a2, b2, block_acc_h2);
            }
            float2 pf = __half22float2(block_acc_h2);
            acc = __fmaf_rn(pf.x + pf.y, scale_ab, acc);
        }
    }

    // Warp reduction
    #pragma unroll
    for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
        acc += __shfl_down_sync(0xffffffff, acc, offset);
    }

    int warp_id = tid_k / WARP_SIZE; // 0..3
    int lane_id = tid_k % WARP_SIZE;

    if (lane_id == 0) {
        smem_acc[warp_id] = acc;
    }
    __syncthreads();

    if (warp_id == 0) {
        float val = (tid_k < 4) ? smem_acc[tid_k] : 0.0f;
        val += __shfl_down_sync(0xffffffff, val, 2);
        val += __shfl_down_sync(0xffffffff, val, 1);
        if (tid_k == 0) {
            size_t c_idx = (size_t)m; // L=0
            C[c_idx] = __float2half(val);
        }
    }
}

// Shared memory B caching kernel (For Large L)
// Copy from v2_merge
__global__ void gemv_kernel_smem(
    const __nv_fp4x2_storage_t* __restrict__ A,
    const __nv_fp4x2_storage_t* __restrict__ B,
    const __nv_fp8_storage_t*  __restrict__ sfa,
    const __nv_fp8_storage_t*  __restrict__ sfb,
    half* __restrict__ C,
    int M, int K, int L
) {
    int m_tile_idx = blockIdx.x;
    int l = blockIdx.y;
    int tid_k = threadIdx.x;   // K-direction
    int tid_m = threadIdx.y;   // M-direction within tile
    int m = m_tile_idx * M_TILE + tid_m;
    if (l >= L) return;
    const int K_sf = K / 16;
    const int K_half = K / 2;
    extern __shared__ char shared_mem[];
    __nv_fp4x2_storage_t* s_B = (__nv_fp4x2_storage_t*)shared_mem;
    __nv_fp8_storage_t* s_sfb = (__nv_fp8_storage_t*)(s_B + K_half);
    const size_t seg_b = (size_t)128 * K_half * l;
    const size_t seg_sfb = (size_t)128 * K_sf * l;
    int total_threads = blockDim.x * blockDim.y;
    int linear_tid = tid_m * blockDim.x + tid_k;
    for (int i = linear_tid; i < K_half / 16; i += total_threads) {
        uint4 b_vec = *reinterpret_cast<const uint4*>(&B[seg_b + i * 16]);
        *reinterpret_cast<uint4*>(&s_B[i * 16]) = b_vec;
    }
    for (int i = (K_half / 16) * 16 + linear_tid; i < K_half; i += total_threads) {
        s_B[i] = B[seg_b + i];
    }
    for (int i = linear_tid; i < K_sf / 16; i += total_threads) {
        uint4 sf_vec = *reinterpret_cast<const uint4*>(&sfb[seg_sfb + i * 16]);
        *reinterpret_cast<uint4*>(&s_sfb[i * 16]) = sf_vec;
    }
    for (int i = (K_sf / 16) * 16 + linear_tid; i < K_sf; i += total_threads) {
        s_sfb[i] = sfb[seg_sfb + i];
    }
    __syncthreads();
    if (m >= M) return;
    const size_t seg_a = (size_t)M * K_half * l + (size_t)m * K_half;
    const size_t seg_sfa = (size_t)M * K_sf * l + (size_t)m * K_sf;
    float acc = 0.0f;
    for (int k_base = tid_k * 2; k_base < K_sf; k_base += SMEM_K_THREADS * 2) {
        __nv_fp8x2_storage_t sfa_pair = *reinterpret_cast<const __nv_fp8x2_storage_t*>(&sfa[seg_sfa + k_base]);
        __nv_fp8x2_storage_t sfb_pair = *reinterpret_cast<const __nv_fp8x2_storage_t*>(&s_sfb[k_base]);
        const __nv_fp8_storage_t* sfa_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&sfa_pair);
        const __nv_fp8_storage_t* sfb_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&sfb_pair);
        uint4 a_vec = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8]);
        uint4 b_vec = *reinterpret_cast<const uint4*>(&s_B[k_base * 8]);
        const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_vec);
        const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_vec);
        #pragma unroll
        for (int i = 0; i < 2; i++) {
            int k = k_base + i;
            if (k >= K_sf) break;
            float scale_a = __half2float(decode_fp8_h(sfa_bytes[i]));
            float scale_b = __half2float(decode_fp8_h(sfb_bytes[i]));
            float scale_ab = scale_a * scale_b;
            float block_sum = 0.0f;
            #pragma unroll
            for (int j = 0; j < 8; j++) {
                __half2 a2 = decode_fp4_h2(a_fp4x2[i * 8 + j]);
                __half2 b2 = decode_fp4_h2(b_fp4x2[i * 8 + j]);
                __half2 prod = __hmul2(a2, b2);
                float2 pf = __half22float2(prod);
                block_sum += (pf.x + pf.y);
            }
            acc = __fmaf_rn(block_sum, scale_ab, acc);
        }
    }
    #pragma unroll
    for (int offset = WARP_SIZE >> 1; offset > 0; offset >>= 1) {
        acc += __shfl_down_sync(0xffffffff, acc, offset);
    }
    if (tid_k == 0) {
        size_t c_idx = (size_t)m + (size_t)l * (size_t)M;
        C[c_idx] = __float2half(acc);
    }
}

// Original L1 kernel (for L=2,3)
__global__ void gemv_kernel_l1(
    const __nv_fp4x2_storage_t* __restrict__ A,
    const __nv_fp4x2_storage_t* __restrict__ B,
    const __nv_fp8_storage_t*  __restrict__ sfa,
    const __nv_fp8_storage_t*  __restrict__ sfb,
    half* __restrict__ C,
    int M, int K, int L
) {
    int m = blockIdx.x;
    int l = blockIdx.y;
    int tid_k = threadIdx.x;
    __shared__ float smem_acc[4];
    if (m >= M || l >= L) return;
    const int K_sf = K / 16;
    const int K_half = K / 2;
    const size_t seg_a = (size_t)M * K_half * l + (size_t)m * K_half;
    const size_t seg_b = (size_t)128 * K_half * l;
    const size_t seg_sfa = (size_t)M * K_sf * l + (size_t)m * K_sf;
    const size_t seg_sfb = (size_t)128 * K_sf * l;
    float acc = 0.0f;
    int k_base = tid_k * 4;
    uint4 a_vec0, a_vec1, b_vec0, b_vec1;
    uint sfa_packed, sfb_packed;
    if (k_base < K_sf) {
        sfa_packed = *reinterpret_cast<const uint*>(&sfa[seg_sfa + k_base]);
        sfb_packed = *reinterpret_cast<const uint*>(&sfb[seg_sfb + k_base]);
        a_vec0 = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8]);
        a_vec1 = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8 + 16]);
        b_vec0 = *reinterpret_cast<const uint4*>(&B[seg_b + k_base * 8]);
        b_vec1 = *reinterpret_cast<const uint4*>(&B[seg_b + k_base * 8 + 16]);
    }
    for (; k_base < K_sf; k_base += K_L1_THREADS * 4) {
        uint4 curr_a0 = a_vec0, curr_a1 = a_vec1;
        uint4 curr_b0 = b_vec0, curr_b1 = b_vec1;
        uint curr_sfa = sfa_packed, curr_sfb = sfb_packed;
        int next_k = k_base + K_L1_THREADS * 4;
        if (next_k < K_sf) {
            sfa_packed = *reinterpret_cast<const uint*>(&sfa[seg_sfa + next_k]);
            sfb_packed = *reinterpret_cast<const uint*>(&sfb[seg_sfb + next_k]);
            a_vec0 = *reinterpret_cast<const uint4*>(&A[seg_a + next_k * 8]);
            a_vec1 = *reinterpret_cast<const uint4*>(&A[seg_a + next_k * 8 + 16]);
            b_vec0 = *reinterpret_cast<const uint4*>(&B[seg_b + next_k * 8]);
            b_vec1 = *reinterpret_cast<const uint4*>(&B[seg_b + next_k * 8 + 16]);
        }
        const __nv_fp8_storage_t* sfa_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&curr_sfa);
        const __nv_fp8_storage_t* sfb_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&curr_sfb);
        const __nv_fp4x2_storage_t* a_fp4x2_0 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_a0);
        const __nv_fp4x2_storage_t* b_fp4x2_0 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_b0);
        #pragma unroll
        for (int i = 0; i < 2; i++) {
            if (k_base + i >= K_sf) break;
            float scale_ab = __half2float(decode_fp8_h(sfa_bytes[i])) * __half2float(decode_fp8_h(sfb_bytes[i]));
            __half2 block_acc_h2 = __float2half2_rn(0.0f);
            #pragma unroll
            for (int j = 0; j < 8; j++) {
                __half2 a2 = decode_fp4_h2(a_fp4x2_0[i * 8 + j]);
                __half2 b2 = decode_fp4_h2(b_fp4x2_0[i * 8 + j]);
                block_acc_h2 = __hfma2(a2, b2, block_acc_h2);
            }
            float2 pf = __half22float2(block_acc_h2);
            acc = __fmaf_rn(pf.x + pf.y, scale_ab, acc);
        }
        const __nv_fp4x2_storage_t* a_fp4x2_1 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_a1);
        const __nv_fp4x2_storage_t* b_fp4x2_1 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_b1);
        #pragma unroll
        for (int i = 0; i < 2; i++) {
            if (k_base + 2 + i >= K_sf) break;
            float scale_ab = __half2float(decode_fp8_h(sfa_bytes[2 + i])) * __half2float(decode_fp8_h(sfb_bytes[2 + i]));
            __half2 block_acc_h2 = __float2half2_rn(0.0f);
            #pragma unroll
            for (int j = 0; j < 8; j++) {
                __half2 a2 = decode_fp4_h2(a_fp4x2_1[i * 8 + j]);
                __half2 b2 = decode_fp4_h2(b_fp4x2_1[i * 8 + j]);
                block_acc_h2 = __hfma2(a2, b2, block_acc_h2);
            }
            float2 pf = __half22float2(block_acc_h2);
            acc = __fmaf_rn(pf.x + pf.y, scale_ab, acc);
        }
    }
    #pragma unroll
    for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
        acc += __shfl_down_sync(0xffffffff, acc, offset);
    }
    int warp_id = tid_k / WARP_SIZE;
    int lane_id = tid_k % WARP_SIZE;
    if (lane_id == 0) smem_acc[warp_id] = acc;
    __syncthreads();
    if (warp_id == 0) {
        float val = (tid_k < 4) ? smem_acc[tid_k] : 0.0f;
        val += __shfl_down_sync(0xffffffff, val, 2);
        val += __shfl_down_sync(0xffffffff, val, 1);
        if (tid_k == 0) {
            size_t c_idx = (size_t)m + (size_t)l * (size_t)M;
            C[c_idx] = __float2half(val);
        }
    }
}


__global__ void gemv_kernel(
    const __nv_fp4x2_storage_t* __restrict__ A,
    const __nv_fp4x2_storage_t* __restrict__ B,
    const __nv_fp8_storage_t*  __restrict__ sfa,
    const __nv_fp8_storage_t*  __restrict__ sfb,
    half* __restrict__ C,
    int M, int K, int L
) {
    int m = blockIdx.x;
    int l_base = blockIdx.y * blockDim.y;
    int tid_k = threadIdx.x;  // K 방향
    int tid_l = threadIdx.y;  // L 방향
    int l = l_base + tid_l;

    if (m >= M || l >= L) return;

    const int K_sf = K / 16;
    const int K_half = K / 2;

    const size_t seg_a = (size_t)M * K_half * l + (size_t)m * K_half;
    const size_t seg_b = (size_t)128 * K_half * l;
    const size_t seg_sfa = (size_t)M * K_sf * l + (size_t)m * K_sf;
    const size_t seg_sfb = (size_t)128 * K_sf * l;

    float acc = 0.0f;

    // 2개 k 블록씩 처리 (원래 로직)
    for (int k_base = tid_k * 2; k_base < K_sf; k_base += K_THREADS * 2) {
        __nv_fp8x2_storage_t sfa_pair = *reinterpret_cast<const __nv_fp8x2_storage_t*>(&sfa[seg_sfa + k_base]);
        __nv_fp8x2_storage_t sfb_pair = *reinterpret_cast<const __nv_fp8x2_storage_t*>(&sfb[seg_sfb + k_base]);

        const __nv_fp8_storage_t* sfa_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&sfa_pair);
        const __nv_fp8_storage_t* sfb_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&sfb_pair);

        uint4 a_vec = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8]);
        uint4 b_vec = *reinterpret_cast<const uint4*>(&B[seg_b + k_base * 8]);

        const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_vec);
        const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_vec);

        #pragma unroll
        for (int i = 0; i < 2; i++) {
            int k = k_base + i;
            if (k >= K_sf) break;

            float scale_a = __half2float(decode_fp8_h(sfa_bytes[i]));
            float scale_b = __half2float(decode_fp8_h(sfb_bytes[i]));
            float scale_ab = scale_a * scale_b;

            float block_sum = 0.0f;

            #pragma unroll
            for (int j = 0; j < 8; j++) {
                __half2 a2 = decode_fp4_h2(a_fp4x2[i * 8 + j]);
                __half2 b2 = decode_fp4_h2(b_fp4x2[i * 8 + j]);

                __half2 prod = __hmul2(a2, b2);
                float2 pf = __half22float2(prod);
                block_sum += (pf.x + pf.y);
            }

            acc = __fmaf_rn(block_sum, scale_ab, acc);
        }
    }

    // Warp reduction
    #pragma unroll
    for (int offset = WARP_SIZE >> 1; offset > 0; offset >>= 1) {
        acc += __shfl_down_sync(0xffffffff, acc, offset);
    }

    size_t c_idx = (size_t)m + (size_t)l * (size_t)M;
    C[c_idx] = __float2half(acc);
}


torch::Tensor gemv_cuda(torch::Tensor A, torch::Tensor B, torch::Tensor C, torch::Tensor sfa, torch::Tensor sfb) {
    int M = A.size(0);
    int K = A.size(1) * 2;
    int L = A.size(2);

    if (L == 1) {
        // Specialized L1 kernel (No L overhead)
        dim3 threadsPerBlock(K_L1_THREADS);
        dim3 blocksPerGrid(M, 1); // L=1
        
        gemv_kernel_l1_specialized<<<blocksPerGrid, threadsPerBlock>>>(
            reinterpret_cast<__nv_fp4x2_storage_t*>(A.data_ptr()),
            reinterpret_cast<__nv_fp4x2_storage_t*>(B.data_ptr()),
            reinterpret_cast<__nv_fp8_storage_t*>(sfa.data_ptr()),
            reinterpret_cast<__nv_fp8_storage_t*>(sfb.data_ptr()),
            reinterpret_cast<half*>(C.data_ptr()),
            M,
            K
        );
    } else if (L < 4) {
        // Use L1 kernel with 4-block pipelining for small L
        dim3 threadsPerBlock(K_L1_THREADS);
        dim3 blocksPerGrid(M, L);

        gemv_kernel_l1<<<blocksPerGrid, threadsPerBlock>>>(
            reinterpret_cast<__nv_fp4x2_storage_t*>(A.data_ptr()),
            reinterpret_cast<__nv_fp4x2_storage_t*>(B.data_ptr()),
            reinterpret_cast<__nv_fp8_storage_t*>(sfa.data_ptr()),
            reinterpret_cast<__nv_fp8_storage_t*>(sfb.data_ptr()),
            reinterpret_cast<half*>(C.data_ptr()),
            M,
            K,
            L
        );
    } else if (L >= 8) {
        // Use shared memory B caching for large L (better B reuse)
        int K_half = K / 2;
        int K_sf = K / 16;
        size_t shared_mem_size = K_half * sizeof(unsigned char) + K_sf * sizeof(unsigned char);

        dim3 threadsPerBlock(SMEM_K_THREADS, M_TILE);
        dim3 blocksPerGrid((M + M_TILE - 1) / M_TILE, L);

        gemv_kernel_smem<<<blocksPerGrid, threadsPerBlock, shared_mem_size>>>(
            reinterpret_cast<__nv_fp4x2_storage_t*>(A.data_ptr()),
            reinterpret_cast<__nv_fp4x2_storage_t*>(B.data_ptr()),
            reinterpret_cast<__nv_fp8_storage_t*>(sfa.data_ptr()),
            reinterpret_cast<__nv_fp8_storage_t*>(sfb.data_ptr()),
            reinterpret_cast<half*>(C.data_ptr()),
            M,
            K,
            L
        );
    } else {
        // L = 4~7: use original kernel
        int L_TILE_SIZE = 4;
        dim3 threadsPerBlock(K_THREADS, L_TILE_SIZE);
        dim3 blocksPerGrid(M, (L + L_TILE_SIZE - 1) / L_TILE_SIZE);

        gemv_kernel<<<blocksPerGrid, threadsPerBlock>>>(
            reinterpret_cast<__nv_fp4x2_storage_t*>(A.data_ptr()),
            reinterpret_cast<__nv_fp4x2_storage_t*>(B.data_ptr()),
            reinterpret_cast<__nv_fp8_storage_t*>(sfa.data_ptr()),
            reinterpret_cast<__nv_fp8_storage_t*>(sfb.data_ptr()),
            reinterpret_cast<half*>(C.data_ptr()),
            M,
            K,
            L
        );
    }

    return C;
}
"""

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

torch::Tensor gemv_cuda(torch::Tensor A, torch::Tensor B, torch::Tensor C, torch::Tensor sfa, torch::Tensor sfb);
"""

gemv_module = load_inline(
    name='gemv_cuda_l1_spec',
    cpp_sources=cpp_code,
    cuda_sources=cuda_kernel_code,
    functions=['gemv_cuda'],
    with_cuda=True,
    extra_cuda_cflags=["-O3", "-use_fast_math", "-gencode=arch=compute_100a,code=sm_100a"],
    verbose=False,
)

def gemv(A, B, C, sfa, sfb):
    if not A.is_cuda or not B.is_cuda or not C.is_cuda or not sfa.is_cuda or not sfb.is_cuda:
        raise RuntimeError("All tensors must be on GPU")
    return gemv_module.gemv_cuda(A, B, C, sfa, sfb)

def custom_kernel(
    data: input_t,
) -> output_t:
    a_ref, b_ref, sfa, sfb, _, _, c_ref = data
    sfa = sfa.to(device=torch.cuda.current_device())
    sfb = sfb.to(device=torch.cuda.current_device())
    return gemv(a_ref, b_ref, c_ref, sfa, sfb)
scrolls · 550 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Best evidence level for this revision: reported

JSON