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
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×fp4fp8
__device__ __forceinline__ __half2 fp8x2_e4m3_to_half2(__nv_fp8x2_e4m3 v) {shared-memory
__shared__ uint4 b_shared1[32];vector-width = float2
float2 x;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 linesreturn *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 1Dint threadID = threadIdx.x;int warpID = threadID / 32;int rowID = warpID;⋯ 8 unchanged linesint 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 linesconst __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 halffor (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 linesauto 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 lineselse {throw std::runtime_error("Unsupported (M, K) combination");}- */cudaError_t err = cudaGetLastError();if (err != cudaSuccess) {⋯ 13 unchanged linestorch::Tensor SFB,torch::Tensor C);"""++extra_cuda_cflags = ["-O3","--use_fast_math",⋯ 24 unchanged linesextra_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