submission 100855
_spatters · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 413 lines, June 9 Researcher Reciprocity License v1.0.
v4a.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-100855?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:676886a76e9b84c5d2d0993edccc1fc8a0c8426a93aaf35e4c22cfc8513346e8
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) {vector-width = uint4
uint4 * a_reg_ptr = reinterpret_cast<uint4 *>(&a_reg_fp4x2[0]);Kernel source
v4a.py413 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<cuda_fp4.h>
#include<cuda_fp16.h>
#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 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(__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, int M_BLOCK, int M_TILE>
__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 rowID, laneID;
get_tile<32>(threadID, rowID, laneID);
int laneOffset = laneID * FP4X2_PER_16B;
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 threadRowIdx = blockRowIdx + rowID;
int batchBlockIdx = blockIdx.z;
int aBatchOffset = 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;
const __nv_fp4x2_e2m1 *gALanePtr = A + aOffset + laneOffset;
const __nv_fp8x2_e4m3 *gSFALanePtr = SFA + sfaOffset + 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]);
__nv_fp8x2_e4m3 sfa_reg_fp8x2;
__nv_fp8x2_e4m3 sfb_reg_fp8x2;
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;
}
}
smol_k += K_BLOCK_SMOL;
}
// 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]);
}
}
}
template<int M, int K, int M_BLOCK, int M_TILE>
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, M_BLOCK, M_TILE><<<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);
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());
// 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);
}
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);
}
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);
}
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);
}
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);
}
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);
}
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);
}
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);
}
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);
}
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);
}
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);
}
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);
}
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",
]
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 · 413 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 100792.
⋯ 13 unchanged lines#include<cuda_fp4.h>#include<cuda_fp16.h>- #define M_BLOCK 4- #define M_TILE 2- #define M_BLOCK_TILED M_BLOCK * M_TILE#define FP4X2_PER_16B 16#define FP8X2_PER_16B 8#define K_BLOCK 32 * FP4X2_PER_16B⋯ 32 unchanged linesreturn *reinterpret_cast<__half*>(&hraw);}- template<int M, int K>+ template<int M, int K, int M_BLOCK, int M_TILE>__global__ void gemv_kernel(const __nv_fp4x2_e2m1* A,const __nv_fp4x2_e2m1* B,⋯ 4 unchanged linesint threadID = threadIdx.x;int rowID, laneID;get_tile<32>(threadID, rowID, laneID);+ int laneOffset = laneID * FP4X2_PER_16B;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;⋯ 6 unchanged linesint threadRowIdx = blockRowIdx + rowID;int batchBlockIdx = blockIdx.z;- int batchOffset = MK * batchBlockIdx;+ int aBatchOffset = 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 = 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;- 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 *gALanePtr = A + aOffset + laneOffset;+ const __nv_fp8x2_e4m3 *gSFALanePtr = SFA + sfaOffset + laneID;- const __nv_fp4x2_e2m1 *gBLanePtr = B + bBatchOffset + FP4X2_PER_16B * 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];- //float2 a_reg_float2[16];- //float2 b_reg_float2[16];__half2 a_reg_half2[16];__half2 b_reg_half2[16];uint4 * a_reg_ptr = reinterpret_cast<uint4 *>(&a_reg_fp4x2[0]);⋯ 2 unchanged lines__nv_fp8x2_e4m3 sfa_reg_fp8x2;__nv_fp8x2_e4m3 sfb_reg_fp8x2;- int laneOffset = laneID * FP4X2_PER_16B;float final_accum[M_TILE] = {0.0f};int smol_k = 0;for (int k_tile=0; k_tile<K; k_tile+=K_BLOCK) {⋯ 2 unchanged lines// read 16B from global to regconst uint4 *gB_ptr = reinterpret_cast<const uint4 *>(gBLanePtr + k_tile);const __nv_fp8x2_e4m3 *gSFB_ptr = (gSFBLanePtr + smol_k);- //const uint4 *gA_ptr = reinterpret_cast<const uint4 *>(gALanePtr + k_tile);- //const __nv_fp8x2_e4m3 *gSFA_ptr = (gSFALanePtr + smol_k);// Read bvals once*b_reg_ptr = *gB_ptr;⋯ 3 unchanged linesb_reg_half2[j] = (fp4x2_e2m1_to_half2(b_reg_fp4x2[j]));}__half2 sfb_vals_h = (fp8x2_e4m3_to_half2(sfb_reg_fp8x2));+// tile over Mfor (int m_tile=0; m_tile<M_TILE; ++m_tile) {int aTileOffset = MBK * m_tile;⋯ 7 unchanged linesfor (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 unrollfor (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);}- __half2 scale = __hmul2(sfa_vals_h, sfb_vals_h);- __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);⋯ 17 unchanged lines}- template<int M, int K>+ template<int M, int K, int M_BLOCK, int M_TILE>void launch_gemv(const __nv_fp4x2_e2m1* A,const __nv_fp4x2_e2m1* B,⋯ 3 unchanged linesdim3 grid,int threads){- gemv_kernel<M, K><<<grid, threads>>>(A, B, SFA, SFB, C);+ gemv_kernel<M, K, M_BLOCK, M_TILE><<<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_TILED), 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());⋯ 3 unchanged lines// K is in units of fp4x2 so half the K of the problem shapesif (M==128 && K==128) {- launch_gemv<128, 128>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==128 && K==768) {- launch_gemv<128, 768>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==128 && K==1536) {- launch_gemv<128, 1536>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==256 && K==3584) {- launch_gemv<256, 3584>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==2432 && K==2304) {- launch_gemv<2432, 2304>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==384 && K==3584) {- launch_gemv<384, 3584>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==512 && K==256) {- launch_gemv<512, 256>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==512 && K==2048) {- launch_gemv<512, 2048>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==512 && K==768) {- launch_gemv<512, 768>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==7168 && K==8192) {- launch_gemv<7168, 8192>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==4096 && K==3584) {- launch_gemv<4096, 3584>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else if (M==7168 && K==1024) {- launch_gemv<7168, 1024>(A_ptr, B_ptr, SFA_ptr, SFB_ptr, C_ptr, grid, threads);+ 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);}else {throw std::runtime_error("Unsupported (M, K) combination");
scrolls · 259 diff lines total
Best evidence level for this revision: reported
JSON