Skip to content
KernelIndex
Search⌘K

submission 95448

_spatters · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-95448?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
36.0µs
#196 of 678
2025-11-22

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:bdb9b787e790be20581f1c1eb63a7c36888af2e737a7e1892ed6e217047b35ab
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) {
shared-memory__shared__ uint4 b_shared1[32];

Kernel source

v3.py420 lines
#!POPCORN leaderboard nvfp4_gemv

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<cuda_fp4.h>
#include<cuda_fp16.h>

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

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

__global__ void debug_print(
        const __nv_fp8x2_e4m3* SFA
        ) {
  __nv_fp8x2_e4m3 sfa_reg_fp8x2;
  float2 x;
  __half2 xh;
  for (int i=0; i<16; i++) {
    sfa_reg_fp8x2 = *(SFA + i);
    xh = fp8x2_e4m3_to_half2(sfa_reg_fp8x2);
    x = __half22float2(xh);
    printf("sfa[%d] %f, sfa[%d] %f \n", i, x.x, i+1, x.y);
  }
}

__global__ void debug_print_scalar(
        const __nv_fp8_e4m3* SFA
        ) {
  __nv_fp8_e4m3 sfa_reg_fp8;
  float x1, x2;
  for (int i=0; i<16; i++) {
    sfa_reg_fp8 = *(SFA + 2*i);
    x1 = __half2float(fp8_e4m3_to_half(sfa_reg_fp8));
    sfa_reg_fp8 = *(SFA + 2*i+1);
    x2 = __half2float(fp8_e4m3_to_half(sfa_reg_fp8));
    printf("sfa[%d] %f, sfa[%d] %f \n", i, x1, i+1, x2);;
  }
}

template<int M, int K>
__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 M,
//		int N,
//		int K,
//		int L
		) {
  // warp layout
  // M/K
  // warp_0
  // warp_1
  // ...
  // warp_BM-1
  // block is 1D
  int threadID = threadIdx.x;
  int warpID = threadID / 32;
  int rowID = warpID;
  int laneID = threadID % 32;

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

  int batchOffset = M * K * batchBlockIdx;
  int bBatchOffset = 128 * K * 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 = M * K * batchBlockIdx / 16;
  int sfbBatchOffset = 128 * K * batchBlockIdx / 16;
  int sfaRowOffset = K * threadRowIdx / 16;
  const unsigned FULL_MASK = 0xffffffff;

  const __nv_fp4x2_e2m1 *gALanePtr = A + batchOffset + rowOffset + FP4X2_PER_16B * laneID;
  const __nv_fp8x2_e4m3 *gSFALanePtr = SFA + sfaBatchOffset + sfaRowOffset + laneID;

  const __nv_fp4x2_e2m1 *gBLanePtr = B + bBatchOffset + FP4X2_PER_16B * laneID;
  const __nv_fp8x2_e4m3 *gSFBLanePtr = SFB + sfbBatchOffset + laneID;

  __nv_fp4x2_e2m1 a_reg_fp4x2[16];
  __nv_fp4x2_e2m1 b_reg_fp4x2[16];
  uint4 * a_reg_ptr = reinterpret_cast<uint4 *>(&a_reg_fp4x2[0]);
  uint4 * b_reg_ptr = reinterpret_cast<uint4 *>(&b_reg_fp4x2[0]);


  __nv_fp8x2_e4m3 sfa_reg_fp8x2;
  __nv_fp8x2_e4m3 sfb_reg_fp8x2;
  __half2 sfa_reg_half2;
  __half2 sfb_reg_half2;
  float final_accum = 0.0f;

  __shared__ uint4 b_shared1[32];
  __shared__ uint4 b_shared2[32];
  __shared__ __nv_fp8x2_e4m3 sfb_shared1[32];
  __shared__ __nv_fp8x2_e4m3 sfb_shared2[32];

  uint4* b_bufs[2] = {b_shared1, b_shared2};
  __nv_fp8x2_e4m3* sfb_bufs[2] = {sfb_shared1, sfb_shared2};
  uint ctr = 0;

  int laneOffset = laneID * FP4X2_PER_16B;
  for (int k_tile=0; k_tile<K; k_tile+=K_BLOCK) {
    int smol_k = k_tile/16;
    bool in_range = laneOffset < K - k_tile;
    if (in_range) {
      if (warpID==0) {
        //const uint4 *gB_ptr = reinterpret_cast<const uint4 *>(B + bBatchOffset + FP4X2_PER_16B*laneID + k_tile);
        //const __nv_fp8x2_e4m3 *gSFB_ptr = (SFB + sfbBatchOffset + laneID + k_tile/16);
        const uint4 *gB_ptr = reinterpret_cast<const uint4 *>(gBLanePtr + k_tile);
        const __nv_fp8x2_e4m3 *gSFB_ptr = (gSFBLanePtr + smol_k);
        //b_shared[laneID] = *gB_ptr;
        //sfb_shared[laneID] = *gSFB_ptr;
        b_bufs[ctr][laneID] = *gB_ptr;
        sfb_bufs[ctr][laneID] = *gSFB_ptr;
      }
    }
    __syncthreads();

    if (in_range) {
    // read 16B from global to reg
    //const uint4 *gA_ptr = reinterpret_cast<const uint4 *>(A + batchOffset + rowOffset + FP4X2_PER_16B*laneID + k_tile);
    //const __nv_fp8x2_e4m3 *gSFA_ptr = (SFA + sfaBatchOffset + sfaRowOffset + laneID + k_tile/16);
    const uint4 *gA_ptr = reinterpret_cast<const uint4 *>(gALanePtr + k_tile);
    const __nv_fp8x2_e4m3 *gSFA_ptr = (gSFALanePtr + smol_k);
    *a_reg_ptr = *gA_ptr;
    sfa_reg_fp8x2 = *gSFA_ptr;

    // TODO: look at coalescing these loads
    //const uint4 *gB_ptr = reinterpret_cast<const uint4 *>(B + bBatchOffset + FP4X2_PER_16B*laneID + k_tile);
    //const __nv_fp8x2_e4m3 *gSFB_ptr = (SFB + sfbBatchOffset + laneID + k_tile/16);
    //*b_reg_ptr = *gB_ptr;
    //sfb_reg_fp8x2 = *gSFB_ptr;
    //*b_reg_ptr = b_shared[laneID];
    //sfb_reg_fp8x2 = sfb_shared[laneID];
    *b_reg_ptr = b_bufs[ctr][laneID];
    sfb_reg_fp8x2 = sfb_bufs[ctr][laneID];

    // a reg is 16B so contains 32 fp4 vals
    // convert fp4x2 to __half2
    sfa_reg_half2 = fp8x2_e4m3_to_half2(sfa_reg_fp8x2);
    sfb_reg_half2 = fp8x2_e4m3_to_half2(sfb_reg_fp8x2);

    //float sfa_vals[2] = {__half2float(__low2half(sfa_reg_half2)), __half2float(__high2half(sfa_reg_half2))};
    //float sfb_vals[2] = {__half2float(__low2half(sfb_reg_half2)), __half2float(__high2half(sfb_reg_half2))};
    float2 sfa_vals = __half22float2(sfa_reg_half2);
    float2 sfb_vals = __half22float2(sfb_reg_half2);
    float scale0 = sfa_vals.x * sfb_vals.x;
    float scale1 = sfa_vals.y * sfb_vals.y;
	float thread_sum = 0.0f;
    #pragma unroll
    for (int j=0; j<8; ++j) {
      float2 a = __half22float2(fp4x2_e2m1_to_half2(a_reg_fp4x2[j]));
      float2 b = __half22float2(fp4x2_e2m1_to_half2(b_reg_fp4x2[j]));
      thread_sum = __fmaf_rn(a.x, b.x, thread_sum);
      thread_sum = __fmaf_rn(a.y, b.y, thread_sum);
    }
    thread_sum *= scale0;
	float thread_sum2 = 0.0f;
    #pragma unroll
    for (int j=8; j<16; ++j) {
      float2 a = __half22float2(fp4x2_e2m1_to_half2(a_reg_fp4x2[j]));
      float2 b = __half22float2(fp4x2_e2m1_to_half2(b_reg_fp4x2[j]));
      //float scale = (j < 8 ? scale0 : scale1);
      //float ax = a.x * scale1;
      //float ay = a.y * scale1;
      thread_sum2= __fmaf_rn(a.x, b.x, thread_sum2);
      thread_sum2 = __fmaf_rn(a.y, b.y, thread_sum2);
    }
    thread_sum = __fmaf_rn(thread_sum2, scale1, thread_sum);
	final_accum += thread_sum;
    }
    //__syncthreads();

  ctr = (ctr + 1) % 2;
  }
  // 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
	// Tree reduction: fold upper half onto lower half
  for (int offset = 16; offset > 0; offset >>= 1) {
    final_accum += __shfl_down_sync(FULL_MASK, final_accum, offset);
  }
  /*
  // now for all threads with laneID = 0 want to write the FP16 val to C
  __shared__ half shared_C[M_BLOCK];
  if (laneID == 0) {
    shared_C[warpID] = __float2half(final_accum);
  }
  __syncthreads();
  // each thread can write 8 FP16 values to global memory in one go
  // we have BLOCK_M FP16 values to write so need BLOCK_M // 8 threads to participate
  if (threadID < M_BLOCK/8) {
    *reinterpret_cast<uint4 *>(C + cOffset + 8*threadID) = *reinterpret_cast<uint4 *>(shared_C + 8*threadID);
  }
  */

  if (laneID == 0) {
    C[cOffset + warpID] = __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)
{
    gemv_kernel<M, K><<<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");
    
    torch::IntArrayRef a_sizes = A.sizes();
    torch::IntArrayRef b_sizes = B.sizes();
    int M = a_sizes[0];
    //int K = a_sizes[1] * 2;
    int K = a_sizes[1];
    int L = a_sizes[2];
    int N = b_sizes[0];

    //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());
    
    /*
    gemv_kernel<<<grid, threads>>>(
            A_ptr,
            B_ptr,
            SFA_ptr,
            SFB_ptr,
            C_ptr,
            M, K, L
        );
        */

    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);
    }
    else if (M==7168 && K==1024) {
      launch_gemv<7168, 1024>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
    }
    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",
    "-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",
]


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,
)

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_module.gemv_cuda(a_ref, b_ref, sfa, sfb, c_ref)
    #torch.cuda.synchronize()
    #print(c_ref)
    return c_ref
scrolls · 420 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 95378.

⋯ 34 unchanged lines
return *reinterpret_cast<__half*>(&hraw);
}
+ __global__ void debug_print(
+ const __nv_fp8x2_e4m3* SFA
+ ) {
+ __nv_fp8x2_e4m3 sfa_reg_fp8x2;
+ float2 x;
+ __half2 xh;
+ for (int i=0; i<16; i++) {
+ sfa_reg_fp8x2 = *(SFA + i);
+ xh = fp8x2_e4m3_to_half2(sfa_reg_fp8x2);
+ x = __half22float2(xh);
+ printf("sfa[%d] %f, sfa[%d] %f \n", i, x.x, i+1, x.y);
+ }
+ }
+
+ __global__ void debug_print_scalar(
+ const __nv_fp8_e4m3* SFA
+ ) {
+ __nv_fp8_e4m3 sfa_reg_fp8;
+ float x1, x2;
+ for (int i=0; i<16; i++) {
+ sfa_reg_fp8 = *(SFA + 2*i);
+ x1 = __half2float(fp8_e4m3_to_half(sfa_reg_fp8));
+ sfa_reg_fp8 = *(SFA + 2*i+1);
+ x2 = __half2float(fp8_e4m3_to_half(sfa_reg_fp8));
+ printf("sfa[%d] %f, sfa[%d] %f \n", i, x1, i+1, x2);;
+ }
+ }
+
+ template<int M, int K>
__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 M,
- int K
+ const __nv_fp8x2_e4m3* SFA,
+ const __nv_fp8x2_e4m3* SFB,
+ half* C
+ // int M,
+ // int N,
+ // int K,
+ // int L
) {
+ // warp layout
+ // M/K
+ // warp_0
+ // warp_1
+ // ...
+ // warp_BM-1
+ // block is 1D
int threadID = threadIdx.x;
int warpID = threadID / 32;
int rowID = warpID;
⋯ 8 unchanged lines
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 = M * K * batchBlockIdx / 16;
int sfbBatchOffset = 128 * K * batchBlockIdx / 16;
int sfaRowOffset = K * threadRowIdx / 16;
⋯ 5 unchanged lines
const __nv_fp4x2_e2m1 *gBLanePtr = B + bBatchOffset + FP4X2_PER_16B * laneID;
const __nv_fp8x2_e4m3 *gSFBLanePtr = SFB + sfbBatchOffset + laneID;
- __nv_fp4x2_e2m1 b_reg_fp4x2[16];
__nv_fp4x2_e2m1 a_reg_fp4x2[16];
- float2 a_reg_float2[16];
- float2 b_reg_float2[16];
+ __nv_fp4x2_e2m1 b_reg_fp4x2[16];
uint4 * a_reg_ptr = reinterpret_cast<uint4 *>(&a_reg_fp4x2[0]);
uint4 * b_reg_ptr = reinterpret_cast<uint4 *>(&b_reg_fp4x2[0]);
+
__nv_fp8x2_e4m3 sfa_reg_fp8x2;
__nv_fp8x2_e4m3 sfb_reg_fp8x2;
+ __half2 sfa_reg_half2;
+ __half2 sfb_reg_half2;
+ float final_accum = 0.0f;
+ __shared__ uint4 b_shared1[32];
+ __shared__ uint4 b_shared2[32];
+ __shared__ __nv_fp8x2_e4m3 sfb_shared1[32];
+ __shared__ __nv_fp8x2_e4m3 sfb_shared2[32];
+
+ uint4* b_bufs[2] = {b_shared1, b_shared2};
+ __nv_fp8x2_e4m3* sfb_bufs[2] = {sfb_shared1, sfb_shared2};
+ uint ctr = 0;
+
int laneOffset = laneID * FP4X2_PER_16B;
- float final_accum = 0.0f;
- int smol_k = 0;
for (int k_tile=0; k_tile<K; k_tile+=K_BLOCK) {
- //int smol_k = k_tile/16;
+ int smol_k = k_tile/16;
bool in_range = laneOffset < K - k_tile;
if (in_range) {
- // read 16B from global to reg
- const uint4 *gA_ptr = reinterpret_cast<const uint4 *>(gALanePtr + k_tile);
- const uint4 *gB_ptr = reinterpret_cast<const uint4 *>(gBLanePtr + k_tile);
- const __nv_fp8x2_e4m3 *gSFA_ptr = (gSFALanePtr + smol_k);
- const __nv_fp8x2_e4m3 *gSFB_ptr = (gSFBLanePtr + smol_k);
-
- // Read vals from global/shared to reg
- *a_reg_ptr = *gA_ptr;
- sfa_reg_fp8x2 = *gSFA_ptr;
- *b_reg_ptr = *gB_ptr;
- sfb_reg_fp8x2 = *gSFB_ptr;
-
- // Convert all a vals to float
- #pragma unroll
- for (int j=0; j<16; ++j) {
- a_reg_float2[j] = __half22float2(fp4x2_e2m1_to_half2(a_reg_fp4x2[j]));
- b_reg_float2[j] = __half22float2(fp4x2_e2m1_to_half2(b_reg_fp4x2[j]));
- //__half2_raw tmp_a = __nv_cvt_fp4x2_to_halfraw2(a_reg_fp4x2[j].__x, __NV_E2M1);
- //__half2_raw tmp_b = __nv_cvt_fp4x2_to_halfraw2(b_reg_fp4x2[j].__x, __NV_E2M1);
- //a_reg_float2[j] = __half22float2(*reinterpret_cast<half2 *>(&tmp_a));
- //b_reg_float2[j] = __half22float2(*reinterpret_cast<half2 *>(&tmp_b));
+ if (warpID==0) {
+ //const uint4 *gB_ptr = reinterpret_cast<const uint4 *>(B + bBatchOffset + FP4X2_PER_16B*laneID + k_tile);
+ //const __nv_fp8x2_e4m3 *gSFB_ptr = (SFB + sfbBatchOffset + laneID + k_tile/16);
+ const uint4 *gB_ptr = reinterpret_cast<const uint4 *>(gBLanePtr + k_tile);
+ const __nv_fp8x2_e4m3 *gSFB_ptr = (gSFBLanePtr + smol_k);
+ //b_shared[laneID] = *gB_ptr;
+ //sfb_shared[laneID] = *gSFB_ptr;
+ b_bufs[ctr][laneID] = *gB_ptr;
+ sfb_bufs[ctr][laneID] = *gSFB_ptr;
}
+ }
+ __syncthreads();
- //__half2_raw tmp_sfa = __nv_cvt_fp8x2_to_halfraw2(sfa_reg_fp8x2.__x, __NV_E4M3);
- //__half2_raw tmp_sfb = __nv_cvt_fp8x2_to_halfraw2(sfb_reg_fp8x2.__x, __NV_E4M3);
- //float2 sfa_vals = __half22float2(*reinterpret_cast<half2 *>(&tmp_sfa));
- //float2 sfb_vals = __half22float2(*reinterpret_cast<half2 *>(&tmp_sfb));
+ if (in_range) {
+ // read 16B from global to reg
+ //const uint4 *gA_ptr = reinterpret_cast<const uint4 *>(A + batchOffset + rowOffset + FP4X2_PER_16B*laneID + k_tile);
+ //const __nv_fp8x2_e4m3 *gSFA_ptr = (SFA + sfaBatchOffset + sfaRowOffset + laneID + k_tile/16);
+ const uint4 *gA_ptr = reinterpret_cast<const uint4 *>(gALanePtr + k_tile);
+ const __nv_fp8x2_e4m3 *gSFA_ptr = (gSFALanePtr + smol_k);
+ *a_reg_ptr = *gA_ptr;
+ sfa_reg_fp8x2 = *gSFA_ptr;
- float2 sfa_vals = __half22float2(fp8x2_e4m3_to_half2(sfa_reg_fp8x2));
- float2 sfb_vals = __half22float2(fp8x2_e4m3_to_half2(sfb_reg_fp8x2));
- float scale0 = sfa_vals.x * sfb_vals.x;
- float scale1 = sfa_vals.y * sfb_vals.y;
+ // TODO: look at coalescing these loads
+ //const uint4 *gB_ptr = reinterpret_cast<const uint4 *>(B + bBatchOffset + FP4X2_PER_16B*laneID + k_tile);
+ //const __nv_fp8x2_e4m3 *gSFB_ptr = (SFB + sfbBatchOffset + laneID + k_tile/16);
+ //*b_reg_ptr = *gB_ptr;
+ //sfb_reg_fp8x2 = *gSFB_ptr;
+ //*b_reg_ptr = b_shared[laneID];
+ //sfb_reg_fp8x2 = sfb_shared[laneID];
+ *b_reg_ptr = b_bufs[ctr][laneID];
+ sfb_reg_fp8x2 = sfb_bufs[ctr][laneID];
- float acc0 = 0.0f;
- float acc1 = 0.0f;
- float* a_reg_float = reinterpret_cast<float *>(a_reg_float2);
- float* b_reg_float = reinterpret_cast<float *>(b_reg_float2);
- #pragma unroll
- for (int j=0; j<16; ++j) {
- acc0 = __fmaf_rn(a_reg_float[j], b_reg_float[j], acc0);
- //acc0 = __fmaf_rn(a_reg_float2[j].x, b_reg_float2[j].x, acc0);
- //acc0 = __fmaf_rn(a_reg_float2[j].y, b_reg_float2[j].y, acc0);
- }
- #pragma unroll
- for (int j=16; j<32; ++j) {
- acc1 = __fmaf_rn(a_reg_float[j], b_reg_float[j], acc1);
- //acc1 = __fmaf_rn(a_reg_float2[j].x, b_reg_float2[j].x, acc1);
- //acc1 = __fmaf_rn(a_reg_float2[j].y, b_reg_float2[j].y, acc1);
- }
- final_accum = __fmaf_rn(acc0, scale0, final_accum);
- final_accum = __fmaf_rn(acc1, scale1, final_accum);
+ // a reg is 16B so contains 32 fp4 vals
+ // convert fp4x2 to __half2
+ sfa_reg_half2 = fp8x2_e4m3_to_half2(sfa_reg_fp8x2);
+ sfb_reg_half2 = fp8x2_e4m3_to_half2(sfb_reg_fp8x2);
+
+ //float sfa_vals[2] = {__half2float(__low2half(sfa_reg_half2)), __half2float(__high2half(sfa_reg_half2))};
+ //float sfb_vals[2] = {__half2float(__low2half(sfb_reg_half2)), __half2float(__high2half(sfb_reg_half2))};
+ float2 sfa_vals = __half22float2(sfa_reg_half2);
+ float2 sfb_vals = __half22float2(sfb_reg_half2);
+ float scale0 = sfa_vals.x * sfb_vals.x;
+ float scale1 = sfa_vals.y * sfb_vals.y;
+ float thread_sum = 0.0f;
+ #pragma unroll
+ for (int j=0; j<8; ++j) {
+ float2 a = __half22float2(fp4x2_e2m1_to_half2(a_reg_fp4x2[j]));
+ float2 b = __half22float2(fp4x2_e2m1_to_half2(b_reg_fp4x2[j]));
+ thread_sum = __fmaf_rn(a.x, b.x, thread_sum);
+ thread_sum = __fmaf_rn(a.y, b.y, thread_sum);
}
- smol_k += 32;
+ thread_sum *= scale0;
+ float thread_sum2 = 0.0f;
+ #pragma unroll
+ for (int j=8; j<16; ++j) {
+ float2 a = __half22float2(fp4x2_e2m1_to_half2(a_reg_fp4x2[j]));
+ float2 b = __half22float2(fp4x2_e2m1_to_half2(b_reg_fp4x2[j]));
+ //float scale = (j < 8 ? scale0 : scale1);
+ //float ax = a.x * scale1;
+ //float ay = a.y * scale1;
+ thread_sum2= __fmaf_rn(a.x, b.x, thread_sum2);
+ thread_sum2 = __fmaf_rn(a.y, b.y, thread_sum2);
+ }
+ thread_sum = __fmaf_rn(thread_sum2, scale1, thread_sum);
+ final_accum += thread_sum;
+ }
+ //__syncthreads();
+
+ ctr = (ctr + 1) % 2;
}
// 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
+ // Tree reduction: fold upper half onto lower half
for (int offset = 16; offset > 0; offset >>= 1) {
final_accum += __shfl_down_sync(FULL_MASK, final_accum, offset);
}
+ /*
+ // now for all threads with laneID = 0 want to write the FP16 val to C
+ __shared__ half shared_C[M_BLOCK];
if (laneID == 0) {
+ shared_C[warpID] = __float2half(final_accum);
+ }
+ __syncthreads();
+ // each thread can write 8 FP16 values to global memory in one go
+ // we have BLOCK_M FP16 values to write so need BLOCK_M // 8 threads to participate
+ if (threadID < M_BLOCK/8) {
+ *reinterpret_cast<uint4 *>(C + cOffset + 8*threadID) = *reinterpret_cast<uint4 *>(shared_C + 8*threadID);
+ }
+ */
+
+ if (laneID == 0) {
C[cOffset + warpID] = __float2half(final_accum);
}
}
-
- /*
template<int M, int K>
void launch_gemv(
const __nv_fp4x2_e2m1* A,
⋯ 6 unchanged lines
{
gemv_kernel<M, K><<<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");
+ 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);
+ torch::IntArrayRef a_sizes = A.sizes();
+ torch::IntArrayRef b_sizes = B.sizes();
+ int M = a_sizes[0];
+ //int K = a_sizes[1] * 2;
+ int K = a_sizes[1];
+ int L = a_sizes[2];
+ int N = b_sizes[0];
//dim3 block(M_BLOCK * 32, 1, 1);
int threads = M_BLOCK * 32;
⋯ 7 unchanged lines
auto SFB_ptr = reinterpret_cast<__nv_fp8x2_e4m3*>(SFB.data_ptr());
auto C_ptr = reinterpret_cast<__half*>(C.data_ptr());
+ /*
gemv_kernel<<<grid, threads>>>(
A_ptr,
B_ptr,
SFA_ptr,
SFB_ptr,
C_ptr,
- M, K
+ M, K, L
);
+ */
- /*
if (M==128 && K==128) {
launch_gemv<128, 128>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);
}
⋯ 33 unchanged lines
else {
throw std::runtime_error("Unsupported (M, K) combination");
}
- */
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
⋯ 13 unchanged lines
torch::Tensor SFB,
torch::Tensor C);
"""
+
+
extra_cuda_cflags = [
"-O3",
"--use_fast_math",
⋯ 24 unchanged lines
extra_cuda_cflags=extra_cuda_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")
⋯ 32 unchanged lines
_, _, 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)
+ gemv_module.gemv_cuda(a_ref, b_ref, sfa, sfb, c_ref)
#torch.cuda.synchronize()
#print(c_ref)
return c_ref
scrolls · 342 diff lines total

Best evidence level for this revision: reported

JSON