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
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×fp4fp8
__device__ __forceinline__ __half2 fp8x2_e4m3_to_half2(__nv_fp8x2_e4m3 v) {stages = 4
constexpr int K_STAGES = 4;vector-width = half2
uint32_t (&out)[16] // 16× half2 bit patternsKernel 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 linessf_vec_size = 16gemv_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 linesreturn *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 lineshalf* 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 sumconstexpr 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 linesdim3 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 linesint 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 shapesif (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