Skip to content
KernelIndex
Search⌘K

submission 107543

_spatters · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

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

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cfd00cef1e2e9adf22964c6f9d2dad85a4d651e652dfa462e35c7205c5472680
license declaredunknown
license concludedunknown
authors_spatters
imported2026-08-15

Techniques

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

fp4__nv_fp4x2_storage_t raw = v.__x; // packed 2×fp4
fp8__device__ __forceinline__ __half2 fp8x2_e4m3_to_half2(__nv_fp8x2_e4m3 v) {
stages = 4constexpr int K_STAGES = 4;
vector-width = half2uint32_t (&out)[16] // 16× half2 bit patterns

Kernel source

v4b.py598 lines
#!POPCORN leaderboard nvfp4_gemv

import os
os.environ["TORCH_CUDA_ARCH_LIST"] = "10.0"

import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# Kernel configuration parameters
sf_vec_size = 16

gemv_cuda_source = r"""
#include<stddef.h>
#include<cuda_fp4.h>
#include<cuda_fp16.h>

#define M_BLOCK 8
#define FP4X2_PER_16B 16
#define FP8X2_PER_16B 8
#define K_BLOCK 512
#define K_BLOCK_SMOL 32
#define ceilDiv(x, y) (((x) + (y) - 1) / (y))


template<int TILE_SIZE>
__device__ __forceinline__
void get_tile(int idx, int& tile_id, int& offset) {
    static_assert((TILE_SIZE & (TILE_SIZE - 1)) == 0, "Must be power of 2");

    constexpr int mask = TILE_SIZE - 1;
    constexpr int shift = __builtin_ctz(TILE_SIZE);

    tile_id = idx >> shift;
    offset  = idx & mask;
}



__device__ __forceinline__
__half2 fp4x2_e2m1_to_half2_ptx(uint16_t raw_bits) {
    uint32_t out_bits;

    asm volatile(
        "{\n"
        "  .reg .b8  b;\n"
        "  .reg .b32 tmp;\n"

        // take low 8 bits = packed fp4x2
        "  mov.b8 b, %1;\n"

        // convert fp4x2 -> f16x2
        "  cvt.rn.f16x2.e2m1x2 %0, b;\n"
        "}\n"
        : "=r"(out_bits)
        : "h"(raw_bits)
    );

    return *reinterpret_cast<__half2*>(&out_bits);
}

__device__ __forceinline__
void convert16_fp4x2_to_half2(
    const uint32_t (&in)[4],    // 4× packed FP4x2 words
    uint32_t (&out)[16]         // 16× half2 bit patterns
) {
    asm volatile(
        "{\n"
        "  .reg .b8 b0, b1, b2, b3;\n"

        // in[0] -> out[0..3]
        "  mov.b32 {b0, b1, b2, b3}, %16;\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"

        // in[1] -> out[4..7]
        "  mov.b32 {b0, b1, b2, b3}, %17;\n"
        "  cvt.rn.f16x2.e2m1x2 %4,  b0;\n"
        "  cvt.rn.f16x2.e2m1x2 %5,  b1;\n"
        "  cvt.rn.f16x2.e2m1x2 %6,  b2;\n"
        "  cvt.rn.f16x2.e2m1x2 %7,  b3;\n"

        // in[2] -> out[8..11]
        "  mov.b32 {b0, b1, b2, b3}, %18;\n"
        "  cvt.rn.f16x2.e2m1x2 %8,  b0;\n"
        "  cvt.rn.f16x2.e2m1x2 %9,  b1;\n"
        "  cvt.rn.f16x2.e2m1x2 %10, b2;\n"
        "  cvt.rn.f16x2.e2m1x2 %11, b3;\n"

        // in[3] -> out[12..15]
        "  mov.b32 {b0, b1, b2, b3}, %19;\n"
        "  cvt.rn.f16x2.e2m1x2 %12, b0;\n"
        "  cvt.rn.f16x2.e2m1x2 %13, b1;\n"
        "  cvt.rn.f16x2.e2m1x2 %14, b2;\n"
        "  cvt.rn.f16x2.e2m1x2 %15, b3;\n"
        "}\n"
        : // 16 outputs: 16 half2 bit patterns
          "=r"(out[0]),  "=r"(out[1]),  "=r"(out[2]),  "=r"(out[3]),
          "=r"(out[4]),  "=r"(out[5]),  "=r"(out[6]),  "=r"(out[7]),
          "=r"(out[8]),  "=r"(out[9]),  "=r"(out[10]), "=r"(out[11]),
          "=r"(out[12]), "=r"(out[13]), "=r"(out[14]), "=r"(out[15])
        : // 4 packed FP4x2 inputs
          "r"(in[0]), "r"(in[1]), "r"(in[2]), "r"(in[3])
    );
}

/*
__device__ __forceinline__
void convert16_fp4x2_to_half2(
    //const __nv_fp4x2_e2m1 (&in)[16],
    const uint16_t (&in_bits)[16],
    uint32_t (&out)[16]
) {
    //const uint16_t* in_bits = reinterpret_cast<const uint16_t*>(in);

    asm volatile(
        "{\n"
        "  cvt.rn.f16x2.e2m1x2 %0,  %16;\n"
        "  cvt.rn.f16x2.e2m1x2 %1,  %17;\n"
        "  cvt.rn.f16x2.e2m1x2 %2,  %18;\n"
        "  cvt.rn.f16x2.e2m1x2 %3,  %19;\n"
        "  cvt.rn.f16x2.e2m1x2 %4,  %20;\n"
        "  cvt.rn.f16x2.e2m1x2 %5,  %21;\n"
        "  cvt.rn.f16x2.e2m1x2 %6,  %22;\n"
        "  cvt.rn.f16x2.e2m1x2 %7,  %23;\n"
        "  cvt.rn.f16x2.e2m1x2 %8,  %24;\n"
        "  cvt.rn.f16x2.e2m1x2 %9,  %25;\n"
        "  cvt.rn.f16x2.e2m1x2 %10, %26;\n"
        "  cvt.rn.f16x2.e2m1x2 %11, %27;\n"
        "  cvt.rn.f16x2.e2m1x2 %12, %28;\n"
        "  cvt.rn.f16x2.e2m1x2 %13, %29;\n"
        "  cvt.rn.f16x2.e2m1x2 %14, %30;\n"
        "  cvt.rn.f16x2.e2m1x2 %15, %31;\n"
        "}\n"
        :
          "=r"(out[0]), "=r"(out[1]), "=r"(out[2]), "=r"(out[3]),
          "=r"(out[4]), "=r"(out[5]), "=r"(out[6]), "=r"(out[7]),
          "=r"(out[8]), "=r"(out[9]), "=r"(out[10]), "=r"(out[11]),
          "=r"(out[12]), "=r"(out[13]), "=r"(out[14]), "=r"(out[15])
        :
          "r"(in_bits[0]&0xFF),  "r"(in_bits[1]&0xFF),  "r"(in_bits[2]&0xFF),  "r"(in_bits[3]&0xFF),
          "r"(in_bits[4]&0xFF),  "r"(in_bits[5]&0xFF),  "r"(in_bits[6]&0xFF),  "r"(in_bits[7]&0xFF),
          "r"(in_bits[8]&0xFF),  "r"(in_bits[9]&0xFF),  "r"(in_bits[10]&0xFF), "r"(in_bits[11]&0xFF),
          "r"(in_bits[12]&0xFF), "r"(in_bits[13]&0xFF), "r"(in_bits[14]&0xFF), "r"(in_bits[15]&0xFF)
    );
}
*/


__device__ __forceinline__ 
uint4 pred_ld_uint4_cs(const uint4* ptr, bool pred) {
    uint4 v;
    v.x = v.y = v.z = v.w = 0u;
    asm volatile(
        "{\n\t"
        "  .reg .pred p;\n\t"
        "  setp.ne.s32 p, %4, 0;             // p = (pred != 0)\n\t"
        "  @p ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%5];\n\t"
        "}\n"
        : "+r"(v.x), "+r"(v.y), "+r"(v.z), "+r"(v.w)
        : "r"((int)pred), "l"(ptr)
    );
    return v;
}

__device__ __forceinline__ 
uint4 pred_ld_uint4_ca(const uint4* ptr, bool pred) {
    uint4 v;
    v.x = v.y = v.z = v.w = 0u;
    asm volatile(
        "{\n\t"
        "  .reg .pred p;\n\t"
        "  setp.ne.s32 p, %4, 0;             // p = (pred != 0)\n\t"
        "  @p ld.global.ca.v4.u32 {%0,%1,%2,%3}, [%5];\n\t"
        "}\n"
        : "+r"(v.x), "+r"(v.y), "+r"(v.z), "+r"(v.w)
        : "r"((int)pred), "l"(ptr)
    );
    return v;
}

__device__ __forceinline__
__half2 fp4x2_e2m1_to_half2_ptx(__nv_fp4x2_e2m1 v) {
    uint16_t raw_bits = *reinterpret_cast<uint16_t*>(&v);
    return fp4x2_e2m1_to_half2_ptx(raw_bits);
}

__device__ __forceinline__ __half2 fp4x2_e2m1_to_half2(__nv_fp4x2_e2m1 v) {
    __nv_fp4x2_storage_t raw = v.__x;  // packed 2×fp4
    __half2_raw hraw = __nv_cvt_fp4x2_to_halfraw2(raw, __NV_E2M1);
    return *reinterpret_cast<__half2*>(&hraw);
}

__device__ __forceinline__ __half2 fp8x2_e4m3_to_half2(__nv_fp8x2_e4m3 v) {
    __nv_fp8x2_storage_t raw = v.__x;
    __half2_raw hraw = __nv_cvt_fp8x2_to_halfraw2(raw, __NV_E4M3);
    return *reinterpret_cast<__half2*>(&hraw);
}

__device__ __forceinline__ __half fp8_e4m3_to_half(__nv_fp8_e4m3 v) {
    __nv_fp8_storage_t raw = v.__x;
    __half_raw hraw = __nv_cvt_fp8_to_halfraw(raw, __NV_E4M3);
    return *reinterpret_cast<__half*>(&hraw);
}


template<int M, int K>
__launch_bounds__(M_BLOCK*32)
__global__ void gemv_kernel(
		const __nv_fp4x2_e2m1* A, 
		const __nv_fp4x2_e2m1* B, 
    const __nv_fp8x2_e4m3* SFA,
    const __nv_fp8x2_e4m3* SFB,
		half* C
		) {
  int threadID = threadIdx.x;
  int warpID, laneID; 
  get_tile<32>(threadID, warpID, laneID);
  int rowID = warpID;
  static_assert(sizeof(__nv_fp4x2_e2m1) == 1, "fp4x2 is not 1 byte");
  static_assert(sizeof(uint4) == 16, "uint4 not 16 bytes");


  constexpr int MK = M * K;
  constexpr int N = 128;
  constexpr int NK = N * K;
  constexpr int MK_SF = MK / 16;
  constexpr int NK_SF = NK / 16;
  constexpr int K_SF = K / 16;

  int blockRowIdx = blockIdx.x * M_BLOCK;
  int threadRowIdx = blockRowIdx + rowID;
  int batchBlockIdx = blockIdx.z;

  int batchOffset = MK * batchBlockIdx;
  int bBatchOffset = NK * batchBlockIdx;
  int rowOffset =  K * threadRowIdx;
  int cOffset = (M * batchBlockIdx + blockRowIdx);

  // scale factor offsets
  // Have K//16 fp8 values per row 
  // We are interpreting the pointer as fp8x2 so we have K//32 values per row
  int sfaBatchOffset = MK_SF * batchBlockIdx;
  int sfbBatchOffset = NK_SF * batchBlockIdx;
  int sfaRowOffset = K_SF * threadRowIdx;

  int laneOffset = laneID * FP4X2_PER_16B;
  const __nv_fp4x2_e2m1 *gALanePtr = A + batchOffset + rowOffset + laneOffset;
  const __nv_fp4x2_e2m1 *gBLanePtr = B + bBatchOffset + laneOffset;
  const uint16_t *gSFALanePtr = reinterpret_cast<const uint16_t *>(SFA + sfaBatchOffset + sfaRowOffset + laneID);
  const uint16_t *gSFBLanePtr = reinterpret_cast<const uint16_t *>(SFB + sfbBatchOffset + laneID);


  constexpr int NUM_TILES = (K + K_BLOCK - 1) / K_BLOCK;
  constexpr int K_STAGES = 4;
  constexpr int K_STAGE_MASK = K_STAGES - 1;
  constexpr int PRELOAD_K = K_STAGES * K_BLOCK;
  constexpr int PRELOAD_K_SMOL = K_STAGES * K_BLOCK_SMOL;
  //__nv_fp4x2_e2m1 a_reg_fp4x2[K_STAGES][16];
  //__nv_fp4x2_e2m1 b_reg_fp4x2[K_STAGES][16];
  uint32_t  a_reg_fp4x2[K_STAGES][4];
  uint32_t  b_reg_fp4x2[K_STAGES][4];

  __nv_fp8x2_e4m3 sfa_reg_fp8x2[K_STAGES];
  __nv_fp8x2_e4m3 sfb_reg_fp8x2[K_STAGES];
  //__half2 a_reg_half2[K_STAGES][16];
  //__half2 b_reg_half2[K_STAGES][16];
  uint32_t a_reg_half2[K_STAGES][16];
  uint32_t b_reg_half2[K_STAGES][16];
  __half2 sfa_vals_h[K_STAGES]; 
  __half2 sfb_vals_h[K_STAGES];
  float final_accum = 0.0f;
  constexpr uint16_t FP8_E4M3_ONE2 = 0x3838;
  constexpr uint4  UINT4_ZERO = uint4{0,0,0,0};
  const __half2 HALF2_ZERO = __float2half2_rn(0.0f);

  const uint4 *gA_ptr, *gB_ptr;
  const uint16_t *gSFA_ptr, *gSFB_ptr;
  // init pointers
  gA_ptr   = reinterpret_cast<const uint4*>(gALanePtr);
  gB_ptr   = reinterpret_cast<const uint4*>(gBLanePtr);
  gSFA_ptr  = gSFALanePtr;
  gSFB_ptr  = gSFBLanePtr;
  
  // Warm up pipeline: prefetch up to K_STAGES tiles
  bool in_range;
  int k_idx = laneOffset;
  #pragma unroll
  for (int stage=0; stage<K_STAGES; ++stage) {
    in_range = k_idx < K;
    //*(reinterpret_cast<uint4 *>(&a_reg_fp4x2[stage][0])) = in_range ? __ldcs(gA_ptr) : UINT4_ZERO;
    *(reinterpret_cast<uint4 *>(&a_reg_fp4x2[stage][0])) = pred_ld_uint4_cs(gA_ptr, in_range); //in_range ? __ldcs(gA_ptr) : UINT4_ZERO;
    *(reinterpret_cast<uint4 *>(&b_reg_fp4x2[stage][0])) = pred_ld_uint4_ca(gB_ptr, in_range);
    *(reinterpret_cast<uint16_t *>(&sfa_reg_fp8x2[stage])) = in_range ? __ldcs(gSFA_ptr) : FP8_E4M3_ONE2;
    *(reinterpret_cast<uint16_t *>(&sfb_reg_fp8x2[stage])) = in_range ? __ldca(gSFB_ptr) : FP8_E4M3_ONE2;
    gA_ptr += 32;
    gB_ptr += 32;
    gSFA_ptr += 32;
    gSFB_ptr += 32;
    k_idx += K_BLOCK;
  }
  // Reset all pointers to what they shold be here (this should be not needed)
  gA_ptr   = reinterpret_cast<const uint4*>(gALanePtr + PRELOAD_K);
  gB_ptr   = reinterpret_cast<const uint4*>(gBLanePtr + PRELOAD_K);
  gSFA_ptr  = gSFALanePtr + PRELOAD_K_SMOL;
  gSFB_ptr  = gSFBLanePtr + PRELOAD_K_SMOL;
  k_idx = laneOffset + PRELOAD_K;
  int stage = 0;
  for (int compute_tile=0;compute_tile<NUM_TILES; ++compute_tile) {
    stage = compute_tile & K_STAGE_MASK;
    // first compute from dis tile
    // DO THE COMPUTE
    sfa_vals_h[stage] = (fp8x2_e4m3_to_half2(sfa_reg_fp8x2[stage]));
    sfb_vals_h[stage] = (fp8x2_e4m3_to_half2(sfb_reg_fp8x2[stage]));
    /*
    #pragma unroll
    for (int j=0; j<FP4X2_PER_16B; ++j) {
      a_reg_half2[stage][j] = (fp4x2_e2m1_to_half2(a_reg_fp4x2[stage][j]));
      b_reg_half2[stage][j] = (fp4x2_e2m1_to_half2(b_reg_fp4x2[stage][j]));
    }
    */
    convert16_fp4x2_to_half2(a_reg_fp4x2[stage], a_reg_half2[stage]);
    convert16_fp4x2_to_half2(b_reg_fp4x2[stage], b_reg_half2[stage]);
    __half2 acc_h0 = HALF2_ZERO;
    __half2 acc_h1 = HALF2_ZERO;
    #pragma unroll
    for (int i = 0; i < 8; ++i) {
      acc_h0 = __hfma2(reinterpret_cast<__half2&>(a_reg_half2[stage][i]), reinterpret_cast<__half2&>(b_reg_half2[stage][i]), acc_h0);
      acc_h1 = __hfma2(reinterpret_cast<__half2&>(a_reg_half2[stage][i+8]), reinterpret_cast<__half2&>(b_reg_half2[stage][i+8]), acc_h1);
      //acc_h0 = __hfma2(a_reg_half2[stage][i], b_reg_half2[stage][i], acc_h0);
      //acc_h1 = __hfma2(a_reg_half2[stage][i+8], b_reg_half2[stage][i+8], acc_h1);
    }
    __half2 scale = __hmul2(sfa_vals_h[stage], sfb_vals_h[stage]);
    __half2 scale0_h = __half2half2(__low2half(scale));
    __half2 scale1_h = __half2half2(__high2half(scale));
    acc_h0 = __hmul2(acc_h0, scale0_h);
    acc_h0 = __hfma2(acc_h1, scale1_h, acc_h0);
    float2 tmp = __half22float2(acc_h0);
    final_accum = final_accum + tmp.x + tmp.y;
    // then load next tile into same slot
    in_range = k_idx < K;
    *(reinterpret_cast<uint4 *>(&a_reg_fp4x2[stage][0])) = pred_ld_uint4_cs(gA_ptr, in_range); 
    *(reinterpret_cast<uint4 *>(&b_reg_fp4x2[stage][0])) = pred_ld_uint4_ca(gB_ptr, in_range);
    *(reinterpret_cast<uint4 *>(&b_reg_fp4x2[stage][0])) = in_range ? __ldca(gB_ptr) : UINT4_ZERO;
    *(reinterpret_cast<uint16_t *>(&sfa_reg_fp8x2[stage])) = in_range ? __ldcs(gSFA_ptr) : FP8_E4M3_ONE2;
    *(reinterpret_cast<uint16_t *>(&sfb_reg_fp8x2[stage])) = in_range ? __ldca(gSFB_ptr) : FP8_E4M3_ONE2;
    // advance da pointers
    gA_ptr += 32;
    gB_ptr += 32;
    gSFA_ptr += 32;
    gSFB_ptr += 32;
    k_idx += K_BLOCK;
  }
  // at this point each thread contains the sum of it's strided values in the row
  // need to use a warp reduction on each warp to compute final row sum
  constexpr unsigned FULL_MASK = 0xffffffff;
  for (int offset = 16; offset > 0; offset >>= 1) {
    final_accum += __shfl_down_sync(FULL_MASK, final_accum, offset);
  }
  if (laneID == 0) {
    C[cOffset + rowID] = __float2half(final_accum);
  }
}


template<int M, int K>
void launch_gemv(
const __nv_fp4x2_e2m1* A,
const __nv_fp4x2_e2m1* B,
const __nv_fp8x2_e4m3* SFA,
const __nv_fp8x2_e4m3* SFB,
half* C,
dim3 grid,
int threads)
{
    auto dis_kernel = gemv_kernel<M, K>;
    /*
    cudaFuncSetAttribute(
        dis_kernel,
        cudaFuncAttributePreferredSharedMemoryCarveout,
        cudaSharedmemCarveoutMaxL1);
        */
    //gemv_kernel<M, K><<<grid, threads>>>(A, B, SFA, SFB, C);
    dis_kernel<<<grid, threads>>>(A, B, SFA, SFB, C);
}


torch::Tensor gemv_cuda(torch::Tensor A, torch::Tensor B, torch::Tensor SFA, torch::Tensor SFB, torch::Tensor C) {
    //TORCH_CHECK(A.device().is_cuda(), "Tensor A must be a CUDA tensor");
    //TORCH_CHECK(B.device().is_cuda(), "Tensor B must be a CUDA tensor");
    //TORCH_CHECK(SFA.device().is_cuda(), "Tensor SFA must be a CUDA tensor");
    //TORCH_CHECK(SFB.device().is_cuda(), "Tensor SFB must be a CUDA tensor");
    //TORCH_CHECK(C.device().is_cuda(), "Tensor C must be a CUDA tensor");
    
    int M = A.size(0); 
    int K = A.size(1); 
    int L = A.size(2); 

    //dim3 block(M_BLOCK * 32, 1, 1);
    int threads = M_BLOCK * 32;
    dim3 grid(ceilDiv(M, M_BLOCK), 1, L);
    //printf("Problem size M: %d, K: %d, N: %d, L: %d \n", M, K, N, L);
    //printf("Threads per block: %d, Block dims (%d, 1, %d)\n", threads, grid.x, grid.z);

    auto A_ptr = reinterpret_cast<__nv_fp4x2_e2m1*>(A.data_ptr());
    auto B_ptr = reinterpret_cast<__nv_fp4x2_e2m1*>(B.data_ptr());
    auto SFA_ptr = reinterpret_cast<__nv_fp8x2_e4m3*>(SFA.data_ptr());
    auto SFB_ptr = reinterpret_cast<__nv_fp8x2_e4m3*>(SFB.data_ptr());
    auto C_ptr = reinterpret_cast<__half*>(C.data_ptr());
  
    // set max l1 for benchmarks
    /*
    if (M==7168 && K==8192) {
      auto dis_kernel = gemv_kernel<7168, 8192>;
      cudaFuncSetAttribute(
          dis_kernel,
          cudaFuncAttributePreferredSharedMemoryCarveout,
          //cudaSharedmemCarveoutMaxL1
          cudaSharedmemCarveoutMaxShared
          );
    }
    else if (M==4096 && K==3584) {
      auto dis_kernel = gemv_kernel<4096, 3384>;
      cudaFuncSetAttribute(
          dis_kernel,
          cudaFuncAttributePreferredSharedMemoryCarveout,
          //cudaSharedmemCarveoutMaxL1);
          ////cudaSharedmemCarveoutMaxL1
          cudaSharedmemCarveoutMaxShared
          );
    }
    else if (M==7168 && K==1024) {
      auto dis_kernel = gemv_kernel<7168, 1024>;
      cudaFuncSetAttribute(
          dis_kernel,
          cudaFuncAttributePreferredSharedMemoryCarveout,
          //cudaSharedmemCarveoutMaxL1);
          //cudaSharedmemCarveoutMaxL1
          cudaSharedmemCarveoutMaxShared
          );
    }
    */
    
    if (M==128 && K==128) {
      launch_gemv<128, 128>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    else if (M==128 && K==768) {
      launch_gemv<128, 768>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    else if (M==128 && K==1536) {
      launch_gemv<128, 1536>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    else if (M==256 && K==3584) {
      launch_gemv<256, 3584>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    else if (M==2432 && K==2304) {
      launch_gemv<2432, 2304>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    else if (M==384 && K==3584) {
      launch_gemv<384, 3584>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    else if (M==512 && K==256) {
      launch_gemv<512, 256>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    else if (M==512 && K==2048) {
      launch_gemv<512, 2048>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    else if (M==512 && K==768) {
      launch_gemv<512, 768>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    else if (M==7168 && K==8192) {
      launch_gemv<7168, 8192>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    else if (M==4096 && K==3584) {
      launch_gemv<4096, 3584>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
      /*
      cudaFuncAttributes attr;
      cudaFuncGetAttributes(&attr, gemv_kernel<4096, 3584>);
      printf("Preferred shared memory carveout = %d\n", attr.preferredShmemCarveout);
      */
    }
    else if (M==7168 && K==1024) {
      launch_gemv<7168, 1024>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
      /*
      cudaFuncAttributes attr;
      cudaFuncGetAttributes(&attr, gemv_kernel<7168, 1024>);
      printf("Preferred shared memory carveout = %d\n", attr.preferredShmemCarveout);
      */
    }
    else {
        throw std::runtime_error("Unsupported (M, K) combination");
    }

    cudaError_t err = cudaGetLastError();
    if (err != cudaSuccess) {
        throw std::runtime_error(cudaGetErrorString(err));
    }
    return C;
}
"""

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

torch::Tensor gemv_cuda(
  torch::Tensor A, 
  torch::Tensor B, 
  torch::Tensor SFA, 
  torch::Tensor SFB, 
  torch::Tensor C);
"""
extra_cuda_cflags = [
    "-O3",
    "--use_fast_math",
    "--fmad=true",
    "--ftz=true",
    "-Xcompiler", "-fno-strict-aliasing",

    # Aggressive math optimizations
    "-Xptxas=-O3",
    #"-Xptxas=--fastmath",

    # Cache behavior
    #"-Xptxas=-dlcm=ca",

    # For debugging performance
    "-Xptxas=--warn-on-spills",
    "-Xptxas=-v",

    # Blackwell target
    #"--gpu-architecture=sm_100a",
    "-gencode=arch=compute_100a,code=sm_100a",
]

extra_cflags = [
    "-O3",
    "-ffast-math",
    "-fno-strict-aliasing",
]


gemv_module = load_inline(
    name='gemv_cuda',
    cpp_sources=gemv_cpp_source,
    cuda_sources=gemv_cuda_source,
    functions=['gemv_cuda'],
    verbose=True,
    extra_cuda_cflags=extra_cuda_cflags,
    extra_cflags=extra_cflags,
)



def gemv_cuda(A, B, SFA, SFB, C):
    if not A.is_cuda or not B.is_cuda or not SFA.is_cuda or not SFB.is_cuda or not C.is_cuda:
        raise RuntimeError("Both tensors must be on GPU")
    return gemv_module.gemv_cuda(A, B, SFA, SFB, C)


# Helper function for ceiling division
def ceil_div(a, b):
    return (a + b - 1) // b


def custom_kernel(
    data: input_t,
) -> output_t:
    """
    PyTorch reference implementation of NVFP4 block-scaled GEMV.
    """
    a_ref, b_ref, sfa, sfb, _, _, c_ref = data
    m, k, l = a_ref.shape
    n, k, l = b_ref.shape
    """
    print(f"K is {k}, n is {n}")
    print(f"A shape {a_ref.shape}")
    print(f"A shape {a_ref.stride()}")
    print(f"SFA shape {sfa.shape}")
    print(f"SFA shape {sfa.stride()}")
    print(f"B shape {b_ref.shape}")
    print(f"B shape {b_ref.stride()}")
    print(f"SFB shape {sfb.shape}")
    print(f"SFB shape {sfb.stride()}")
    print(f"C shape {c_ref.shape}")
    print(f"C shape {c_ref.stride()}")
    """

    # Get dimensions from MxNxL layout
    _, _, l = c_ref.shape
    #print(sfa.shape, sfa.stride())
    #print(f"SFA[0,0:32,0]: {sfa[0,:32,0].reshape(-1,2)}")
    gemv_cuda(a_ref, b_ref, sfa, sfb, c_ref)
    #torch.cuda.synchronize()
    #print(c_ref)
    return c_ref
scrolls · 598 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 100855.

⋯ 10 unchanged lines
sf_vec_size = 16
gemv_cuda_source = r"""
+ #include<stddef.h>
#include<cuda_fp4.h>
#include<cuda_fp16.h>
+ #define M_BLOCK 8
#define FP4X2_PER_16B 16
#define FP8X2_PER_16B 8
- #define K_BLOCK 32 * FP4X2_PER_16B
- #define K_BLOCK_SMOL 32 * FP4X2_PER_16B / 16
+ #define K_BLOCK 512
+ #define K_BLOCK_SMOL 32
#define ceilDiv(x, y) (((x) + (y) - 1) / (y))
⋯ 10 unchanged lines
}
+
+ __device__ __forceinline__
+ __half2 fp4x2_e2m1_to_half2_ptx(uint16_t raw_bits) {
+ uint32_t out_bits;
+
+ asm volatile(
+ "{\n"
+ " .reg .b8 b;\n"
+ " .reg .b32 tmp;\n"
+
+ // take low 8 bits = packed fp4x2
+ " mov.b8 b, %1;\n"
+
+ // convert fp4x2 -> f16x2
+ " cvt.rn.f16x2.e2m1x2 %0, b;\n"
+ "}\n"
+ : "=r"(out_bits)
+ : "h"(raw_bits)
+ );
+
+ return *reinterpret_cast<__half2*>(&out_bits);
+ }
+
+ __device__ __forceinline__
+ void convert16_fp4x2_to_half2(
+ const uint32_t (&in)[4], // 4× packed FP4x2 words
+ uint32_t (&out)[16] // 16× half2 bit patterns
+ ) {
+ asm volatile(
+ "{\n"
+ " .reg .b8 b0, b1, b2, b3;\n"
+
+ // in[0] -> out[0..3]
+ " mov.b32 {b0, b1, b2, b3}, %16;\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"
+
+ // in[1] -> out[4..7]
+ " mov.b32 {b0, b1, b2, b3}, %17;\n"
+ " cvt.rn.f16x2.e2m1x2 %4, b0;\n"
+ " cvt.rn.f16x2.e2m1x2 %5, b1;\n"
+ " cvt.rn.f16x2.e2m1x2 %6, b2;\n"
+ " cvt.rn.f16x2.e2m1x2 %7, b3;\n"
+
+ // in[2] -> out[8..11]
+ " mov.b32 {b0, b1, b2, b3}, %18;\n"
+ " cvt.rn.f16x2.e2m1x2 %8, b0;\n"
+ " cvt.rn.f16x2.e2m1x2 %9, b1;\n"
+ " cvt.rn.f16x2.e2m1x2 %10, b2;\n"
+ " cvt.rn.f16x2.e2m1x2 %11, b3;\n"
+
+ // in[3] -> out[12..15]
+ " mov.b32 {b0, b1, b2, b3}, %19;\n"
+ " cvt.rn.f16x2.e2m1x2 %12, b0;\n"
+ " cvt.rn.f16x2.e2m1x2 %13, b1;\n"
+ " cvt.rn.f16x2.e2m1x2 %14, b2;\n"
+ " cvt.rn.f16x2.e2m1x2 %15, b3;\n"
+ "}\n"
+ : // 16 outputs: 16 half2 bit patterns
+ "=r"(out[0]), "=r"(out[1]), "=r"(out[2]), "=r"(out[3]),
+ "=r"(out[4]), "=r"(out[5]), "=r"(out[6]), "=r"(out[7]),
+ "=r"(out[8]), "=r"(out[9]), "=r"(out[10]), "=r"(out[11]),
+ "=r"(out[12]), "=r"(out[13]), "=r"(out[14]), "=r"(out[15])
+ : // 4 packed FP4x2 inputs
+ "r"(in[0]), "r"(in[1]), "r"(in[2]), "r"(in[3])
+ );
+ }
+
+ /*
+ __device__ __forceinline__
+ void convert16_fp4x2_to_half2(
+ //const __nv_fp4x2_e2m1 (&in)[16],
+ const uint16_t (&in_bits)[16],
+ uint32_t (&out)[16]
+ ) {
+ //const uint16_t* in_bits = reinterpret_cast<const uint16_t*>(in);
+
+ asm volatile(
+ "{\n"
+ " cvt.rn.f16x2.e2m1x2 %0, %16;\n"
+ " cvt.rn.f16x2.e2m1x2 %1, %17;\n"
+ " cvt.rn.f16x2.e2m1x2 %2, %18;\n"
+ " cvt.rn.f16x2.e2m1x2 %3, %19;\n"
+ " cvt.rn.f16x2.e2m1x2 %4, %20;\n"
+ " cvt.rn.f16x2.e2m1x2 %5, %21;\n"
+ " cvt.rn.f16x2.e2m1x2 %6, %22;\n"
+ " cvt.rn.f16x2.e2m1x2 %7, %23;\n"
+ " cvt.rn.f16x2.e2m1x2 %8, %24;\n"
+ " cvt.rn.f16x2.e2m1x2 %9, %25;\n"
+ " cvt.rn.f16x2.e2m1x2 %10, %26;\n"
+ " cvt.rn.f16x2.e2m1x2 %11, %27;\n"
+ " cvt.rn.f16x2.e2m1x2 %12, %28;\n"
+ " cvt.rn.f16x2.e2m1x2 %13, %29;\n"
+ " cvt.rn.f16x2.e2m1x2 %14, %30;\n"
+ " cvt.rn.f16x2.e2m1x2 %15, %31;\n"
+ "}\n"
+ :
+ "=r"(out[0]), "=r"(out[1]), "=r"(out[2]), "=r"(out[3]),
+ "=r"(out[4]), "=r"(out[5]), "=r"(out[6]), "=r"(out[7]),
+ "=r"(out[8]), "=r"(out[9]), "=r"(out[10]), "=r"(out[11]),
+ "=r"(out[12]), "=r"(out[13]), "=r"(out[14]), "=r"(out[15])
+ :
+ "r"(in_bits[0]&0xFF), "r"(in_bits[1]&0xFF), "r"(in_bits[2]&0xFF), "r"(in_bits[3]&0xFF),
+ "r"(in_bits[4]&0xFF), "r"(in_bits[5]&0xFF), "r"(in_bits[6]&0xFF), "r"(in_bits[7]&0xFF),
+ "r"(in_bits[8]&0xFF), "r"(in_bits[9]&0xFF), "r"(in_bits[10]&0xFF), "r"(in_bits[11]&0xFF),
+ "r"(in_bits[12]&0xFF), "r"(in_bits[13]&0xFF), "r"(in_bits[14]&0xFF), "r"(in_bits[15]&0xFF)
+ );
+ }
+ */
+
+
+ __device__ __forceinline__
+ uint4 pred_ld_uint4_cs(const uint4* ptr, bool pred) {
+ uint4 v;
+ v.x = v.y = v.z = v.w = 0u;
+ asm volatile(
+ "{\n\t"
+ " .reg .pred p;\n\t"
+ " setp.ne.s32 p, %4, 0; // p = (pred != 0)\n\t"
+ " @p ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%5];\n\t"
+ "}\n"
+ : "+r"(v.x), "+r"(v.y), "+r"(v.z), "+r"(v.w)
+ : "r"((int)pred), "l"(ptr)
+ );
+ return v;
+ }
+
+ __device__ __forceinline__
+ uint4 pred_ld_uint4_ca(const uint4* ptr, bool pred) {
+ uint4 v;
+ v.x = v.y = v.z = v.w = 0u;
+ asm volatile(
+ "{\n\t"
+ " .reg .pred p;\n\t"
+ " setp.ne.s32 p, %4, 0; // p = (pred != 0)\n\t"
+ " @p ld.global.ca.v4.u32 {%0,%1,%2,%3}, [%5];\n\t"
+ "}\n"
+ : "+r"(v.x), "+r"(v.y), "+r"(v.z), "+r"(v.w)
+ : "r"((int)pred), "l"(ptr)
+ );
+ return v;
+ }
+
+ __device__ __forceinline__
+ __half2 fp4x2_e2m1_to_half2_ptx(__nv_fp4x2_e2m1 v) {
+ uint16_t raw_bits = *reinterpret_cast<uint16_t*>(&v);
+ return fp4x2_e2m1_to_half2_ptx(raw_bits);
+ }
+
__device__ __forceinline__ __half2 fp4x2_e2m1_to_half2(__nv_fp4x2_e2m1 v) {
__nv_fp4x2_storage_t raw = v.__x; // packed 2×fp4
__half2_raw hraw = __nv_cvt_fp4x2_to_halfraw2(raw, __NV_E2M1);
⋯ 12 unchanged lines
return *reinterpret_cast<__half*>(&hraw);
}
- template<int M, int K, int M_BLOCK, int M_TILE>
+
+ template<int M, int K>
+ __launch_bounds__(M_BLOCK*32)
__global__ void gemv_kernel(
const __nv_fp4x2_e2m1* A,
const __nv_fp4x2_e2m1* B,
⋯ 2 unchanged lines
half* C
) {
int threadID = threadIdx.x;
- int rowID, laneID;
- get_tile<32>(threadID, rowID, laneID);
- int laneOffset = laneID * FP4X2_PER_16B;
+ int warpID, laneID;
+ get_tile<32>(threadID, warpID, laneID);
+ int rowID = warpID;
+ static_assert(sizeof(__nv_fp4x2_e2m1) == 1, "fp4x2 is not 1 byte");
+ static_assert(sizeof(uint4) == 16, "uint4 not 16 bytes");
+
constexpr int MK = M * K;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
constexpr int N = 128;
constexpr int NK = N * K;
constexpr int MK_SF = MK / 16;
constexpr int NK_SF = NK / 16;
constexpr int K_SF = K / 16;
- constexpr int MBK = M_BLOCK * K;
- constexpr int MBK_SF = M_BLOCK * K_SF;
- int blockRowIdx = blockIdx.x * M_BLOCK_TILED;
+ int blockRowIdx = blockIdx.x * M_BLOCK;
int threadRowIdx = blockRowIdx + rowID;
int batchBlockIdx = blockIdx.z;
- int aBatchOffset = MK * batchBlockIdx;
+ int batchOffset = MK * batchBlockIdx;
int bBatchOffset = NK * batchBlockIdx;
int rowOffset = K * threadRowIdx;
- int aOffset = aBatchOffset + rowOffset;
int cOffset = (M * batchBlockIdx + blockRowIdx);
// scale factor offsets
// Have K//16 fp8 values per row
// We are interpreting the pointer as fp8x2 so we have K//32 values per row
- //int sfaBatchOffset = MK_SF * batchBlockIdx;
- //int sfbBatchOffset = NK_SF * batchBlockIdx;
- //int sfaRowOffset = K_SF * threadRowIdx;
- //int sfaBatchOffset = aBatchOffset >> 4;
- //int sfaRowOffset = rowOffset >> 4;
- int sfaOffset = aOffset >> 4;
- int sfbBatchOffset = bBatchOffset >> 4;
+ int sfaBatchOffset = MK_SF * batchBlockIdx;
+ int sfbBatchOffset = NK_SF * batchBlockIdx;
+ int sfaRowOffset = K_SF * threadRowIdx;
- const __nv_fp4x2_e2m1 *gALanePtr = A + aOffset + laneOffset;
- const __nv_fp8x2_e4m3 *gSFALanePtr = SFA + sfaOffset + laneID;
+ int laneOffset = laneID * FP4X2_PER_16B;
+ const __nv_fp4x2_e2m1 *gALanePtr = A + batchOffset + rowOffset + laneOffset;
+ const __nv_fp4x2_e2m1 *gBLanePtr = B + bBatchOffset + laneOffset;
+ const uint16_t *gSFALanePtr = reinterpret_cast<const uint16_t *>(SFA + sfaBatchOffset + sfaRowOffset + laneID);
+ const uint16_t *gSFBLanePtr = reinterpret_cast<const uint16_t *>(SFB + sfbBatchOffset + laneID);
- const __nv_fp4x2_e2m1 *gBLanePtr = B + bBatchOffset + laneOffset;
- const __nv_fp8x2_e4m3 *gSFBLanePtr = SFB + sfbBatchOffset + laneID;
- __nv_fp4x2_e2m1 b_reg_fp4x2[16];
- __nv_fp4x2_e2m1 a_reg_fp4x2[16];
- __half2 a_reg_half2[16];
- __half2 b_reg_half2[16];
- uint4 * a_reg_ptr = reinterpret_cast<uint4 *>(&a_reg_fp4x2[0]);
- uint4 * b_reg_ptr = reinterpret_cast<uint4 *>(&b_reg_fp4x2[0]);
+ constexpr int NUM_TILES = (K + K_BLOCK - 1) / K_BLOCK;
+ constexpr int K_STAGES = 4;
+ constexpr int K_STAGE_MASK = K_STAGES - 1;
+ constexpr int PRELOAD_K = K_STAGES * K_BLOCK;
+ constexpr int PRELOAD_K_SMOL = K_STAGES * K_BLOCK_SMOL;
+ //__nv_fp4x2_e2m1 a_reg_fp4x2[K_STAGES][16];
+ //__nv_fp4x2_e2m1 b_reg_fp4x2[K_STAGES][16];
+ uint32_t a_reg_fp4x2[K_STAGES][4];
+ uint32_t b_reg_fp4x2[K_STAGES][4];
- __nv_fp8x2_e4m3 sfa_reg_fp8x2;
- __nv_fp8x2_e4m3 sfb_reg_fp8x2;
+ __nv_fp8x2_e4m3 sfa_reg_fp8x2[K_STAGES];
+ __nv_fp8x2_e4m3 sfb_reg_fp8x2[K_STAGES];
+ //__half2 a_reg_half2[K_STAGES][16];
+ //__half2 b_reg_half2[K_STAGES][16];
+ uint32_t a_reg_half2[K_STAGES][16];
+ uint32_t b_reg_half2[K_STAGES][16];
+ __half2 sfa_vals_h[K_STAGES];
+ __half2 sfb_vals_h[K_STAGES];
+ float final_accum = 0.0f;
+ constexpr uint16_t FP8_E4M3_ONE2 = 0x3838;
+ constexpr uint4 UINT4_ZERO = uint4{0,0,0,0};
+ const __half2 HALF2_ZERO = __float2half2_rn(0.0f);
- float final_accum[M_TILE] = {0.0f};
- int smol_k = 0;
- for (int k_tile=0; k_tile<K; k_tile+=K_BLOCK) {
- bool in_range = laneOffset < K - k_tile;
- if (in_range) {
- // read 16B from global to reg
- const uint4 *gB_ptr = reinterpret_cast<const uint4 *>(gBLanePtr + k_tile);
- const __nv_fp8x2_e4m3 *gSFB_ptr = (gSFBLanePtr + smol_k);
-
- // Read bvals once
- *b_reg_ptr = *gB_ptr;
- sfb_reg_fp8x2 = *gSFB_ptr;
- #pragma unroll
- for (int j=0; j<16; ++j) {
- b_reg_half2[j] = (fp4x2_e2m1_to_half2(b_reg_fp4x2[j]));
- }
- __half2 sfb_vals_h = (fp8x2_e4m3_to_half2(sfb_reg_fp8x2));
-
- // tile over M
- for (int m_tile=0; m_tile<M_TILE; ++m_tile) {
- int aTileOffset = MBK * m_tile;
- int sfaTileOffset = MBK_SF * m_tile;
- const uint4 *gA_ptr = reinterpret_cast<const uint4 *>(gALanePtr + aTileOffset + k_tile);
- const __nv_fp8x2_e4m3 *gSFA_ptr = (gSFALanePtr + sfaTileOffset + smol_k);
- *a_reg_ptr = *gA_ptr;
- sfa_reg_fp8x2 = *gSFA_ptr;
- __half2 sfa_vals_h = (fp8x2_e4m3_to_half2(sfa_reg_fp8x2));
- #pragma unroll
- for (int j=0; j<16; ++j) {
- a_reg_half2[j] = (fp4x2_e2m1_to_half2(a_reg_fp4x2[j]));
- }
- __half2 scale = __hmul2(sfa_vals_h, sfb_vals_h);
- __half2 acc_h0 = __float2half2_rn(0.0f);
- __half2 acc_h1 = __float2half2_rn(0.0f);
- __half2 scale0_h = __half2half2(__low2half(scale));
- __half2 scale1_h = __half2half2(__high2half(scale));
- #pragma unroll
- for (int i = 0; i < 8; ++i) {
- acc_h0 = __hfma2(a_reg_half2[i], b_reg_half2[i], acc_h0);
- acc_h1 = __hfma2(a_reg_half2[i+8], b_reg_half2[i+8], acc_h1);
- }
- acc_h0 = __hmul2(acc_h0, scale0_h);
- acc_h0 = __hfma2(acc_h1, scale1_h, acc_h0);
- float2 tmp = __half22float2(acc_h0);
- final_accum[m_tile] = final_accum[m_tile] + tmp.x + tmp.y;
- }
+ const uint4 *gA_ptr, *gB_ptr;
+ const uint16_t *gSFA_ptr, *gSFB_ptr;
+ // init pointers
+ gA_ptr = reinterpret_cast<const uint4*>(gALanePtr);
+ gB_ptr = reinterpret_cast<const uint4*>(gBLanePtr);
+ gSFA_ptr = gSFALanePtr;
+ gSFB_ptr = gSFBLanePtr;
+
+ // Warm up pipeline: prefetch up to K_STAGES tiles
+ bool in_range;
+ int k_idx = laneOffset;
+ #pragma unroll
+ for (int stage=0; stage<K_STAGES; ++stage) {
+ in_range = k_idx < K;
+ //*(reinterpret_cast<uint4 *>(&a_reg_fp4x2[stage][0])) = in_range ? __ldcs(gA_ptr) : UINT4_ZERO;
+ *(reinterpret_cast<uint4 *>(&a_reg_fp4x2[stage][0])) = pred_ld_uint4_cs(gA_ptr, in_range); //in_range ? __ldcs(gA_ptr) : UINT4_ZERO;
+ *(reinterpret_cast<uint4 *>(&b_reg_fp4x2[stage][0])) = pred_ld_uint4_ca(gB_ptr, in_range);
+ *(reinterpret_cast<uint16_t *>(&sfa_reg_fp8x2[stage])) = in_range ? __ldcs(gSFA_ptr) : FP8_E4M3_ONE2;
+ *(reinterpret_cast<uint16_t *>(&sfb_reg_fp8x2[stage])) = in_range ? __ldca(gSFB_ptr) : FP8_E4M3_ONE2;
+ gA_ptr += 32;
+ gB_ptr += 32;
+ gSFA_ptr += 32;
+ gSFB_ptr += 32;
+ k_idx += K_BLOCK;
+ }
+ // Reset all pointers to what they shold be here (this should be not needed)
+ gA_ptr = reinterpret_cast<const uint4*>(gALanePtr + PRELOAD_K);
+ gB_ptr = reinterpret_cast<const uint4*>(gBLanePtr + PRELOAD_K);
+ gSFA_ptr = gSFALanePtr + PRELOAD_K_SMOL;
+ gSFB_ptr = gSFBLanePtr + PRELOAD_K_SMOL;
+ k_idx = laneOffset + PRELOAD_K;
+ int stage = 0;
+ for (int compute_tile=0;compute_tile<NUM_TILES; ++compute_tile) {
+ stage = compute_tile & K_STAGE_MASK;
+ // first compute from dis tile
+ // DO THE COMPUTE
+ sfa_vals_h[stage] = (fp8x2_e4m3_to_half2(sfa_reg_fp8x2[stage]));
+ sfb_vals_h[stage] = (fp8x2_e4m3_to_half2(sfb_reg_fp8x2[stage]));
+ /*
+ #pragma unroll
+ for (int j=0; j<FP4X2_PER_16B; ++j) {
+ a_reg_half2[stage][j] = (fp4x2_e2m1_to_half2(a_reg_fp4x2[stage][j]));
+ b_reg_half2[stage][j] = (fp4x2_e2m1_to_half2(b_reg_fp4x2[stage][j]));
}
- smol_k += K_BLOCK_SMOL;
+ */
+ convert16_fp4x2_to_half2(a_reg_fp4x2[stage], a_reg_half2[stage]);
+ convert16_fp4x2_to_half2(b_reg_fp4x2[stage], b_reg_half2[stage]);
+ __half2 acc_h0 = HALF2_ZERO;
+ __half2 acc_h1 = HALF2_ZERO;
+ #pragma unroll
+ for (int i = 0; i < 8; ++i) {
+ acc_h0 = __hfma2(reinterpret_cast<__half2&>(a_reg_half2[stage][i]), reinterpret_cast<__half2&>(b_reg_half2[stage][i]), acc_h0);
+ acc_h1 = __hfma2(reinterpret_cast<__half2&>(a_reg_half2[stage][i+8]), reinterpret_cast<__half2&>(b_reg_half2[stage][i+8]), acc_h1);
+ //acc_h0 = __hfma2(a_reg_half2[stage][i], b_reg_half2[stage][i], acc_h0);
+ //acc_h1 = __hfma2(a_reg_half2[stage][i+8], b_reg_half2[stage][i+8], acc_h1);
+ }
+ __half2 scale = __hmul2(sfa_vals_h[stage], sfb_vals_h[stage]);
+ __half2 scale0_h = __half2half2(__low2half(scale));
+ __half2 scale1_h = __half2half2(__high2half(scale));
+ acc_h0 = __hmul2(acc_h0, scale0_h);
+ acc_h0 = __hfma2(acc_h1, scale1_h, acc_h0);
+ float2 tmp = __half22float2(acc_h0);
+ final_accum = final_accum + tmp.x + tmp.y;
+ // then load next tile into same slot
+ in_range = k_idx < K;
+ *(reinterpret_cast<uint4 *>(&a_reg_fp4x2[stage][0])) = pred_ld_uint4_cs(gA_ptr, in_range);
+ *(reinterpret_cast<uint4 *>(&b_reg_fp4x2[stage][0])) = pred_ld_uint4_ca(gB_ptr, in_range);
+ *(reinterpret_cast<uint4 *>(&b_reg_fp4x2[stage][0])) = in_range ? __ldca(gB_ptr) : UINT4_ZERO;
+ *(reinterpret_cast<uint16_t *>(&sfa_reg_fp8x2[stage])) = in_range ? __ldcs(gSFA_ptr) : FP8_E4M3_ONE2;
+ *(reinterpret_cast<uint16_t *>(&sfb_reg_fp8x2[stage])) = in_range ? __ldca(gSFB_ptr) : FP8_E4M3_ONE2;
+ // advance da pointers
+ gA_ptr += 32;
+ gB_ptr += 32;
+ gSFA_ptr += 32;
+ gSFB_ptr += 32;
+ k_idx += K_BLOCK;
}
// at this point each thread contains the sum of it's strided values in the row
// need to use a warp reduction on each warp to compute final row sum
constexpr unsigned FULL_MASK = 0xffffffff;
-
- for (int m_tile=0; m_tile<M_TILE; ++m_tile) {
- for (int offset = 16; offset > 0; offset >>= 1) {
- final_accum[m_tile] += __shfl_down_sync(FULL_MASK, final_accum[m_tile], offset);
- }
- if (laneID == 0) {
- C[cOffset + m_tile*M_BLOCK + rowID] = __float2half(final_accum[m_tile]);
- }
+ for (int offset = 16; offset > 0; offset >>= 1) {
+ final_accum += __shfl_down_sync(FULL_MASK, final_accum, offset);
}
+ if (laneID == 0) {
+ C[cOffset + rowID] = __float2half(final_accum);
+ }
}
- template<int M, int K, int M_BLOCK, int M_TILE>
+ template<int M, int K>
void launch_gemv(
const __nv_fp4x2_e2m1* A,
const __nv_fp4x2_e2m1* B,
⋯ 3 unchanged lines
dim3 grid,
int threads)
{
- gemv_kernel<M, K, M_BLOCK, M_TILE><<<grid, threads>>>(A, B, SFA, SFB, C);
+ auto dis_kernel = gemv_kernel<M, K>;
+ /*
+ cudaFuncSetAttribute(
+ dis_kernel,
+ cudaFuncAttributePreferredSharedMemoryCarveout,
+ cudaSharedmemCarveoutMaxL1);
+ */
+ //gemv_kernel<M, K><<<grid, threads>>>(A, B, SFA, SFB, C);
+ dis_kernel<<<grid, threads>>>(A, B, SFA, SFB, C);
}
⋯ 8 unchanged lines
int K = A.size(1);
int L = A.size(2);
+ //dim3 block(M_BLOCK * 32, 1, 1);
+ int threads = M_BLOCK * 32;
+ dim3 grid(ceilDiv(M, M_BLOCK), 1, L);
+ //printf("Problem size M: %d, K: %d, N: %d, L: %d \n", M, K, N, L);
+ //printf("Threads per block: %d, Block dims (%d, 1, %d)\n", threads, grid.x, grid.z);
auto A_ptr = reinterpret_cast<__nv_fp4x2_e2m1*>(A.data_ptr());
auto B_ptr = reinterpret_cast<__nv_fp4x2_e2m1*>(B.data_ptr());
auto SFA_ptr = reinterpret_cast<__nv_fp8x2_e4m3*>(SFA.data_ptr());
auto SFB_ptr = reinterpret_cast<__nv_fp8x2_e4m3*>(SFB.data_ptr());
auto C_ptr = reinterpret_cast<__half*>(C.data_ptr());
+
+ // set max l1 for benchmarks
+ /*
+ if (M==7168 && K==8192) {
+ auto dis_kernel = gemv_kernel<7168, 8192>;
+ cudaFuncSetAttribute(
+ dis_kernel,
+ cudaFuncAttributePreferredSharedMemoryCarveout,
+ //cudaSharedmemCarveoutMaxL1
+ cudaSharedmemCarveoutMaxShared
+ );
+ }
+ else if (M==4096 && K==3584) {
+ auto dis_kernel = gemv_kernel<4096, 3384>;
+ cudaFuncSetAttribute(
+ dis_kernel,
+ cudaFuncAttributePreferredSharedMemoryCarveout,
+ //cudaSharedmemCarveoutMaxL1);
+ ////cudaSharedmemCarveoutMaxL1
+ cudaSharedmemCarveoutMaxShared
+ );
+ }
+ else if (M==7168 && K==1024) {
+ auto dis_kernel = gemv_kernel<7168, 1024>;
+ cudaFuncSetAttribute(
+ dis_kernel,
+ cudaFuncAttributePreferredSharedMemoryCarveout,
+ //cudaSharedmemCarveoutMaxL1);
+ //cudaSharedmemCarveoutMaxL1
+ cudaSharedmemCarveoutMaxShared
+ );
+ }
+ */
- // K is in units of fp4x2 so half the K of the problem shapes
if (M==128 && K==128) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 4;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<128, 128, 2, 4>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<128, 128>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
else if (M==128 && K==768) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 4;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<128, 768, 2, 4>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<128, 768>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
else if (M==128 && K==1536) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 4;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<128, 1536, 2, 4>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<128, 1536>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
else if (M==256 && K==3584) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 4;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<256, 3584, 2, 4>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<256, 3584>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
else if (M==2432 && K==2304) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 4;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<2432, 2304, 2, 4>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<2432, 2304>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
else if (M==384 && K==3584) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 4;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<384, 3584, 2, 4>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<384, 3584>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
else if (M==512 && K==256) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 4;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<512, 256, 2, 4>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<512, 256>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
else if (M==512 && K==2048) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 4;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<512, 2048, 2, 4>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<512, 2048>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
else if (M==512 && K==768) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 4;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<512, 768, 2, 4>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<512, 768>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
else if (M==7168 && K==8192) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 2;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<7168, 8192, M_BLOCK, M_TILE>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<7168, 8192>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
else if (M==4096 && K==3584) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 2;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<4096, 3584, M_BLOCK, M_TILE>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<4096, 3584>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ /*
+ cudaFuncAttributes attr;
+ cudaFuncGetAttributes(&attr, gemv_kernel<4096, 3584>);
+ printf("Preferred shared memory carveout = %d\n", attr.preferredShmemCarveout);
+ */
}
else if (M==7168 && K==1024) {
- constexpr int M_BLOCK = 2;
- constexpr int M_TILE = 4;
- constexpr int M_BLOCK_TILED = M_BLOCK * M_TILE;
- int threads = M_BLOCK * 32;
- dim3 grid(ceilDiv(M, M_BLOCK_TILED), 1, L);
- launch_gemv<7168, 1024, 2, 4>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ launch_gemv<7168, 1024>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
+ /*
+ cudaFuncAttributes attr;
+ cudaFuncGetAttributes(&attr, gemv_kernel<7168, 1024>);
+ printf("Preferred shared memory carveout = %d\n", attr.preferredShmemCarveout);
+ */
}
else {
throw std::runtime_error("Unsupported (M, K) combination");
⋯ 29 unchanged lines
#"-Xptxas=--fastmath",
# Cache behavior
- "-Xptxas=-dlcm=ca",
+ #"-Xptxas=-dlcm=ca",
# For debugging performance
"-Xptxas=--warn-on-spills",
"-Xptxas=-v",
# Blackwell target
- "--gpu-architecture=sm_100a",
+ #"--gpu-architecture=sm_100a",
+ "-gencode=arch=compute_100a,code=sm_100a",
]
extra_cflags = [
scrolls · 633 diff lines total

Best evidence level for this revision: reported

JSON