submission 110022
macto · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 528 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-110022?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:e99d109a4fe5130319486dc5b601cbdcf246f9144453ed39de2358621bb7897d
license declaredunknown
license concludedunknown
authorsmacto
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" \shared-memory
__shared__ __align__(16) uint8_t sh_a[NUM_STAGES][ROWS_PER_BLOCK][K_TILE_BYTES];vector-width = half2
__device__ __forceinline__ void cvt_fp4_to_f16(half2* out, uint32_t in) {Kernel source
submission.py528 lines
"""
Dual-module hybrid: Loads direct and pipelined kernels as SEPARATE modules
to avoid any compiler optimization interference between them.
"""
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# =============================================================================
# DIRECT KERNEL MODULE (exact copy from v13)
# =============================================================================
DIRECT_CUDA = r'''
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
enum class CacheHint { CS, CA };
template<CacheHint hint>
__device__ __forceinline__ void load_vec4(uint32_t* dst, const void* src) {
if constexpr (hint == CacheHint::CS) {
asm volatile("ld.global.cs.v4.u32 {%0, %1, %2, %3}, [%4];"
: "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
} else {
asm volatile("ld.global.ca.v4.u32 {%0, %1, %2, %3}, [%4];"
: "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
}
}
template<CacheHint hint>
__device__ __forceinline__ void load_u16(uint16_t* dst, const void* src) {
if constexpr (hint == CacheHint::CS) {
asm volatile("ld.global.cs.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
} else {
asm volatile("ld.global.ca.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
}
}
__device__ __forceinline__ void cvt_fp4_to_f16(half2* out, uint32_t in) {
asm volatile(
"{\n\t"
".reg .b8 b0, b1, b2, b3;\n\t"
"mov.b32 {b0, b1, b2, b3}, %4;\n\t"
"cvt.rn.f16x2.e2m1x2 %0, b0;\n\t"
"cvt.rn.f16x2.e2m1x2 %1, b1;\n\t"
"cvt.rn.f16x2.e2m1x2 %2, b2;\n\t"
"cvt.rn.f16x2.e2m1x2 %3, b3;\n\t"
"}"
: "=r"(reinterpret_cast<int&>(out[0])),
"=r"(reinterpret_cast<int&>(out[1])),
"=r"(reinterpret_cast<int&>(out[2])),
"=r"(reinterpret_cast<int&>(out[3]))
: "r"(in)
);
}
__device__ __forceinline__ half2 cvt_fp8_to_f16(uint16_t in) {
int out;
asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(in));
return *reinterpret_cast<half2*>(&out);
}
__device__ __forceinline__ float dot_scaled_32fp4(
const uint32_t* a_data,
const uint32_t* b_data,
half2 scale
) {
half2 a_f16[16], b_f16[16];
#pragma unroll
for (int i = 0; i < 4; i++) {
cvt_fp4_to_f16(a_f16 + i * 4, a_data[i]);
cvt_fp4_to_f16(b_f16 + i * 4, b_data[i]);
}
half2 acc0 = __hmul2(a_f16[0], b_f16[0]);
half2 acc1 = __hmul2(a_f16[8], b_f16[8]);
#pragma unroll
for (int i = 1; i < 8; i++) {
acc0 = __hfma2(a_f16[i], b_f16[i], acc0);
acc1 = __hfma2(a_f16[8 + i], b_f16[8 + i], acc1);
}
half sum0 = __hadd(acc0.x, acc0.y);
half sum1 = __hadd(acc1.x, acc1.y);
float result = 0.0f;
asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"
: "+f"(result) : "h"(*(uint16_t*)&sum0), "h"(*(uint16_t*)&scale.x));
asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"
: "+f"(result) : "h"(*(uint16_t*)&sum1), "h"(*(uint16_t*)&scale.y));
return result;
}
template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CacheHint A_HINT, CacheHint B_HINT>
__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)
void gemv_kernel(
const uint8_t* __restrict__ a,
const uint8_t* __restrict__ b,
const uint8_t* __restrict__ sfa,
const uint8_t* __restrict__ sfb,
half* __restrict__ c,
int M, int K, int L, int N_pad
) {
const int row_base = blockIdx.x * ROWS_PER_BLOCK;
const int batch = blockIdx.y;
const int tid = threadIdx.x;
const int warp_id = tid / THREADS_PER_ROW;
const int lane = tid % THREADS_PER_ROW;
const int row = row_base + warp_id;
if (row >= M) return;
const int K_bytes = K / 2;
const int K_sf = K / 16;
const size_t a_off = (size_t)batch * M * K_bytes + row * K_bytes;
const size_t b_off = (size_t)batch * N_pad * K_bytes;
const size_t sfa_off = (size_t)batch * M * K_sf + row * K_sf;
const size_t sfb_off = (size_t)batch * N_pad * K_sf;
const uint8_t* a_ptr = a + a_off;
const uint8_t* b_ptr = b + b_off;
const uint8_t* sfa_ptr = sfa + sfa_off;
const uint8_t* sfb_ptr = sfb + sfb_off;
float acc = 0.0f;
const int num_sf_pairs = K_sf / 2;
#pragma unroll 4
for (int sp = lane; sp < num_sf_pairs; sp += THREADS_PER_ROW) {
const int sf_idx = sp * 2;
const int byte_idx = sf_idx * 8;
uint16_t sfa_packed, sfb_packed;
load_u16<A_HINT>(&sfa_packed, sfa_ptr + sf_idx);
load_u16<B_HINT>(&sfb_packed, sfb_ptr + sf_idx);
half2 scale = __hmul2(cvt_fp8_to_f16(sfa_packed), cvt_fp8_to_f16(sfb_packed));
uint32_t a_regs[4], b_regs[4];
load_vec4<A_HINT>(a_regs, a_ptr + byte_idx);
load_vec4<B_HINT>(b_regs, b_ptr + byte_idx);
acc += dot_scaled_32fp4(a_regs, b_regs, scale);
}
#pragma unroll
for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {
acc += __shfl_down_sync(0xffffffff, acc, offset);
}
if (lane == 0) {
c[(size_t)row + (size_t)batch * M] = __float2half(acc);
}
}
template<int ROWS, int THREADS, CacheHint A_HINT, CacheHint B_HINT>
void dispatch(torch::Tensor& c, torch::Tensor& a, torch::Tensor& b,
torch::Tensor& sfa, torch::Tensor& sfb, int m, int k, int l, int n) {
dim3 grid((m + ROWS - 1) / ROWS, l);
dim3 block(ROWS * THREADS);
gemv_kernel<ROWS, THREADS, A_HINT, B_HINT><<<grid, block>>>(
a.data_ptr<uint8_t>(), b.data_ptr<uint8_t>(),
sfa.data_ptr<uint8_t>(), sfb.data_ptr<uint8_t>(),
reinterpret_cast<half*>(c.data_ptr<at::Half>()), m, k, l, n);
}
void run_direct(
torch::Tensor C, torch::Tensor A, torch::Tensor B,
torch::Tensor SFA, torch::Tensor SFB,
int m, int k, int l, int n_pad
) {
constexpr auto CS = CacheHint::CS;
constexpr auto CA = CacheHint::CA;
const int k_bytes = k / 2;
const bool b_fits_l1 = (k_bytes <= 32 * 1024);
if (k >= 8192) {
if (b_fits_l1) {
dispatch<8, 32, CS, CA>(C, A, B, SFA, SFB, m, k, l, n_pad);
} else {
dispatch<8, 32, CS, CS>(C, A, B, SFA, SFB, m, k, l, n_pad);
}
} else if (k >= 4096) {
dispatch<4, 32, CS, CA>(C, A, B, SFA, SFB, m, k, l, n_pad);
} else {
dispatch<8, 16, CS, CA>(C, A, B, SFA, SFB, m, k, l, n_pad);
}
}
'''
DIRECT_CPP = r'''
#include <torch/extension.h>
void run_direct(torch::Tensor C, torch::Tensor A, torch::Tensor B,
torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad);
'''
# =============================================================================
# PIPELINED KERNEL MODULE (exact copy from v14)
# =============================================================================
PIPELINED_CUDA = r'''
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#define ASYNC_CP_16(dst, src) \
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" \
:: "r"((uint32_t)__cvta_generic_to_shared(dst)), "l"(src))
#define ASYNC_CP_4(dst, src) \
asm volatile("cp.async.ca.shared.global [%0], [%1], 4;" \
:: "r"((uint32_t)__cvta_generic_to_shared(dst)), "l"(src))
#define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")
#define ASYNC_WAIT(n) asm volatile("cp.async.wait_group %0;" :: "n"(n))
#define NUM_STAGES 2
__device__ __forceinline__ void cvt_fp4_to_f16(half2* out, uint32_t in) {
asm volatile(
"{\n\t"
".reg .b8 b0, b1, b2, b3;\n\t"
"mov.b32 {b0, b1, b2, b3}, %4;\n\t"
"cvt.rn.f16x2.e2m1x2 %0, b0;\n\t"
"cvt.rn.f16x2.e2m1x2 %1, b1;\n\t"
"cvt.rn.f16x2.e2m1x2 %2, b2;\n\t"
"cvt.rn.f16x2.e2m1x2 %3, b3;\n\t"
"}"
: "=r"(reinterpret_cast<int&>(out[0])),
"=r"(reinterpret_cast<int&>(out[1])),
"=r"(reinterpret_cast<int&>(out[2])),
"=r"(reinterpret_cast<int&>(out[3]))
: "r"(in)
);
}
__device__ __forceinline__ half2 cvt_fp8_to_f16(uint16_t in) {
int out;
asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(in));
return *reinterpret_cast<half2*>(&out);
}
__device__ __forceinline__ float dot_scaled_32fp4(
const uint8_t* a_smem,
const uint8_t* b_smem,
half2 scale
) {
const uint32_t* a_data = reinterpret_cast<const uint32_t*>(a_smem);
const uint32_t* b_data = reinterpret_cast<const uint32_t*>(b_smem);
half2 a_f16[16], b_f16[16];
#pragma unroll
for (int i = 0; i < 4; i++) {
cvt_fp4_to_f16(a_f16 + i * 4, a_data[i]);
cvt_fp4_to_f16(b_f16 + i * 4, b_data[i]);
}
half2 acc0 = __hmul2(a_f16[0], b_f16[0]);
half2 acc1 = __hmul2(a_f16[8], b_f16[8]);
#pragma unroll
for (int i = 1; i < 8; i++) {
acc0 = __hfma2(a_f16[i], b_f16[i], acc0);
acc1 = __hfma2(a_f16[8 + i], b_f16[8 + i], acc1);
}
half sum0 = __hadd(acc0.x, acc0.y);
half sum1 = __hadd(acc1.x, acc1.y);
float result = 0.0f;
asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"
: "+f"(result) : "h"(*(uint16_t*)&sum0), "h"(*(uint16_t*)&scale.x));
asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"
: "+f"(result) : "h"(*(uint16_t*)&sum1), "h"(*(uint16_t*)&scale.y));
return result;
}
template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, int K_TILE>
__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)
void gemv_pipelined_kernel(
const uint8_t* __restrict__ a,
const uint8_t* __restrict__ b,
const uint8_t* __restrict__ sfa,
const uint8_t* __restrict__ sfb,
half* __restrict__ c,
int M, int K, int L, int N_pad
) {
constexpr int K_TILE_BYTES = K_TILE / 2;
constexpr int K_TILE_SF = K_TILE / 16;
__shared__ __align__(16) uint8_t sh_a[NUM_STAGES][ROWS_PER_BLOCK][K_TILE_BYTES];
__shared__ __align__(16) uint8_t sh_b[NUM_STAGES][K_TILE_BYTES];
__shared__ __align__(16) uint8_t sh_sfa[NUM_STAGES][ROWS_PER_BLOCK][K_TILE_SF];
__shared__ __align__(16) uint8_t sh_sfb[NUM_STAGES][K_TILE_SF];
const int row_base = blockIdx.x * ROWS_PER_BLOCK;
const int batch = blockIdx.y;
const int tid = threadIdx.x;
const int warp_id = tid / THREADS_PER_ROW;
const int lane = tid % THREADS_PER_ROW;
const int row = row_base + warp_id;
if (row >= M) return;
const int K_bytes = K / 2;
const int K_sf = K / 16;
const int num_tiles = (K_bytes + K_TILE_BYTES - 1) / K_TILE_BYTES;
const size_t a_base = (size_t)batch * M * K_bytes;
const size_t b_base = (size_t)batch * N_pad * K_bytes;
const size_t sfa_base = (size_t)batch * M * K_sf;
const size_t sfb_base = (size_t)batch * N_pad * K_sf;
auto issue_tile = [&](int tile, int stage) {
const int k_off_bytes = tile * K_TILE_BYTES;
const int k_off_sf = tile * K_TILE_SF;
const int bytes_this = min(K_TILE_BYTES, K_bytes - k_off_bytes);
const int sf_this = min(K_TILE_SF, K_sf - k_off_sf);
if (warp_id < ROWS_PER_BLOCK && (row_base + warp_id) < M) {
const uint8_t* a_row = a + a_base + (row_base + warp_id) * K_bytes + k_off_bytes;
for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16) {
if (i + 16 <= bytes_this) {
ASYNC_CP_16(&sh_a[stage][warp_id][i], a_row + i);
}
}
const uint8_t* sfa_row = sfa + sfa_base + (row_base + warp_id) * K_sf + k_off_sf;
for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4) {
if (i + 4 <= sf_this) {
ASYNC_CP_4(&sh_sfa[stage][warp_id][i], sfa_row + i);
}
}
}
if (warp_id == 0) {
const uint8_t* b_ptr = b + b_base + k_off_bytes;
for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16) {
if (i + 16 <= bytes_this) {
ASYNC_CP_16(&sh_b[stage][i], b_ptr + i);
}
}
const uint8_t* sfb_ptr = sfb + sfb_base + k_off_sf;
for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4) {
if (i + 4 <= sf_this) {
ASYNC_CP_4(&sh_sfb[stage][i], sfb_ptr + i);
}
}
}
ASYNC_COMMIT();
};
issue_tile(0, 0);
ASYNC_WAIT(0);
__syncthreads();
float acc = 0.0f;
for (int tile = 0; tile < num_tiles; tile++) {
const int stage = tile % NUM_STAGES;
const int next_stage = (tile + 1) % NUM_STAGES;
if (tile + 1 < num_tiles) {
issue_tile(tile + 1, next_stage);
}
const int k_off_sf = tile * K_TILE_SF;
const int sf_this = min(K_TILE_SF, K_sf - k_off_sf);
const int num_sf_pairs = sf_this / 2;
#pragma unroll 2
for (int sp = lane; sp < num_sf_pairs; sp += THREADS_PER_ROW) {
const int sf_idx = sp * 2;
const int byte_idx = sf_idx * 8;
uint16_t sfa_packed = *reinterpret_cast<const uint16_t*>(&sh_sfa[stage][warp_id][sf_idx]);
uint16_t sfb_packed = *reinterpret_cast<const uint16_t*>(&sh_sfb[stage][sf_idx]);
half2 scale = __hmul2(cvt_fp8_to_f16(sfa_packed), cvt_fp8_to_f16(sfb_packed));
acc += dot_scaled_32fp4(
&sh_a[stage][warp_id][byte_idx],
&sh_b[stage][byte_idx],
scale
);
}
if (tile + 1 < num_tiles) {
ASYNC_WAIT(0);
__syncthreads();
}
}
#pragma unroll
for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {
acc += __shfl_down_sync(0xffffffff, acc, offset);
}
if (lane == 0) {
c[(size_t)row + (size_t)batch * M] = __float2half(acc);
}
}
template<int ROWS, int THREADS, int K_TILE>
void dispatch(torch::Tensor& c, torch::Tensor& a, torch::Tensor& b,
torch::Tensor& sfa, torch::Tensor& sfb, int m, int k, int l, int n) {
dim3 grid((m + ROWS - 1) / ROWS, l);
dim3 block(ROWS * THREADS);
gemv_pipelined_kernel<ROWS, THREADS, K_TILE><<<grid, block>>>(
a.data_ptr<uint8_t>(), b.data_ptr<uint8_t>(),
sfa.data_ptr<uint8_t>(), sfb.data_ptr<uint8_t>(),
reinterpret_cast<half*>(c.data_ptr<at::Half>()), m, k, l, n);
}
void run_pipelined(
torch::Tensor C, torch::Tensor A, torch::Tensor B,
torch::Tensor SFA, torch::Tensor SFB,
int m, int k, int l, int n_pad, int k_tile_hint
) {
if (k >= 8192) {
// Try different K_TILE values for large K
// k_tile_hint: 0=2048, 1=1024, 2=4096, 3=8192
switch (k_tile_hint) {
case 1: dispatch<4, 32, 1024>(C, A, B, SFA, SFB, m, k, l, n_pad); break;
case 2: dispatch<4, 32, 4096>(C, A, B, SFA, SFB, m, k, l, n_pad); break;
case 3: dispatch<4, 32, 8192>(C, A, B, SFA, SFB, m, k, l, n_pad); break;
default: dispatch<4, 32, 2048>(C, A, B, SFA, SFB, m, k, l, n_pad); break;
}
} else if (k >= 4096) {
dispatch<4, 32, 1024>(C, A, B, SFA, SFB, m, k, l, n_pad);
} else {
dispatch<8, 16, 2048>(C, A, B, SFA, SFB, m, k, l, n_pad);
}
}
'''
PIPELINED_CPP = r'''
#include <torch/extension.h>
void run_pipelined(torch::Tensor C, torch::Tensor A, torch::Tensor B,
torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad, int k_tile_hint);
'''
# =============================================================================
# Module loading
# =============================================================================
# =============================================================================
# Configuration
# =============================================================================
# K_TILE options for k >= 8192:
# 0 = 2048 (default)
# 1 = 1024
# 2 = 4096
# 3 = 8192
K_TILE_HINT = 0 # <-- Change this to test different K_TILE values
_direct_module = None
_pipelined_module = None
def _load_direct():
global _direct_module
if _direct_module is None:
_direct_module = load_inline(
name='nvfp4_gemv_direct_v16',
cpp_sources=DIRECT_CPP,
cuda_sources=DIRECT_CUDA,
functions=['run_direct'],
extra_cuda_cflags=[
'-O3', '--use_fast_math', '-std=c++17',
'-gencode=arch=compute_100a,code=sm_100a',
],
verbose=False,
)
return _direct_module
def _load_pipelined():
global _pipelined_module
if _pipelined_module is None:
_pipelined_module = load_inline(
name='nvfp4_gemv_pipelined_v16b', # Updated to force recompile
cpp_sources=PIPELINED_CPP,
cuda_sources=PIPELINED_CUDA,
functions=['run_pipelined'],
extra_cuda_cflags=[
'-O3', '--use_fast_math', '-std=c++17',
'-gencode=arch=compute_100a,code=sm_100a',
],
verbose=False,
)
return _pipelined_module
def custom_kernel(data: input_t) -> output_t:
a, b, sfa, sfb, _, _, c = data
m, k_packed, l = a.shape
k = k_packed * 2
n_pad = b.shape[0]
if not sfa.is_cuda:
sfa = sfa.to(a.device)
if not sfb.is_cuda:
sfb = sfb.to(a.device)
c_out = c.squeeze(1)
a_u8 = a.view(torch.uint8)
b_u8 = b.view(torch.uint8)
sfa_u8 = sfa.view(torch.uint8)
sfb_u8 = sfb.view(torch.uint8)
# Selection heuristic based on benchmarks:
# - k=16384, l=1: pipelined faster (24.6 vs 30.8)
# - k=7168, l=8: direct faster (38.5 vs 51.0)
# - k=2048, l=4: same (16.4)
use_pipelined = (k >= 8192 and l <= 2) or (k <= 4096 and l >= 4)
if use_pipelined:
_load_pipelined().run_pipelined(c_out, a_u8, b_u8, sfa_u8, sfb_u8, m, k, l, n_pad, 2)
else:
_load_direct().run_direct(c_out, a_u8, b_u8, sfa_u8, sfb_u8, m, k, l, n_pad)
return c
scrolls · 528 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 109980.
+ """+ Dual-module hybrid: Loads direct and pipelined kernels as SEPARATE modules+ to avoid any compiler optimization interference between them.+ """import torchfrom torch.utils.cpp_extension import load_inlinefrom task import input_t, output_t- # Combine both v13 (direct) and v14 (pipelined) kernels with selection logic- CUDA_SRC = r'''+ # =============================================================================+ # DIRECT KERNEL MODULE (exact copy from v13)+ # =============================================================================++ DIRECT_CUDA = r'''#include <torch/extension.h>#include <cuda_fp16.h>#include <cuda_fp4.h>#include <cuda_fp8.h>- // =============================================================================- // Common Utilities- // =============================================================================-enum class CacheHint { CS, CA };template<CacheHint hint>⋯ 40 unchanged linesreturn *reinterpret_cast<half2*>(&out);}- // =============================================================================- // DIRECT KERNEL (from v13)- // =============================================================================-- __device__ __forceinline__ float dot_scaled_direct(+ __device__ __forceinline__ float dot_scaled_32fp4(const uint32_t* a_data,const uint32_t* b_data,half2 scale⋯ 28 unchanged linestemplate<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CacheHint A_HINT, CacheHint B_HINT>__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)- void gemv_direct_kernel(+ void gemv_kernel(const uint8_t* __restrict__ a,const uint8_t* __restrict__ b,const uint8_t* __restrict__ sfa,⋯ 41 unchanged linesload_vec4<A_HINT>(a_regs, a_ptr + byte_idx);load_vec4<B_HINT>(b_regs, b_ptr + byte_idx);- acc += dot_scaled_direct(a_regs, b_regs, scale);+ acc += dot_scaled_32fp4(a_regs, b_regs, scale);}#pragma unroll⋯ 6 unchanged lines}}- // =============================================================================- // PIPELINED KERNEL (from v14)- // =============================================================================+ template<int ROWS, int THREADS, CacheHint A_HINT, CacheHint B_HINT>+ void dispatch(torch::Tensor& c, torch::Tensor& a, torch::Tensor& b,+ torch::Tensor& sfa, torch::Tensor& sfb, int m, int k, int l, int n) {+ dim3 grid((m + ROWS - 1) / ROWS, l);+ dim3 block(ROWS * THREADS);+ gemv_kernel<ROWS, THREADS, A_HINT, B_HINT><<<grid, block>>>(+ a.data_ptr<uint8_t>(), b.data_ptr<uint8_t>(),+ sfa.data_ptr<uint8_t>(), sfb.data_ptr<uint8_t>(),+ reinterpret_cast<half*>(c.data_ptr<at::Half>()), m, k, l, n);+ }+ void run_direct(+ torch::Tensor C, torch::Tensor A, torch::Tensor B,+ torch::Tensor SFA, torch::Tensor SFB,+ int m, int k, int l, int n_pad+ ) {+ constexpr auto CS = CacheHint::CS;+ constexpr auto CA = CacheHint::CA;++ const int k_bytes = k / 2;+ const bool b_fits_l1 = (k_bytes <= 32 * 1024);++ if (k >= 8192) {+ if (b_fits_l1) {+ dispatch<8, 32, CS, CA>(C, A, B, SFA, SFB, m, k, l, n_pad);+ } else {+ dispatch<8, 32, CS, CS>(C, A, B, SFA, SFB, m, k, l, n_pad);+ }+ } else if (k >= 4096) {+ dispatch<4, 32, CS, CA>(C, A, B, SFA, SFB, m, k, l, n_pad);+ } else {+ dispatch<8, 16, CS, CA>(C, A, B, SFA, SFB, m, k, l, n_pad);+ }+ }+ '''++ DIRECT_CPP = r'''+ #include <torch/extension.h>+ void run_direct(torch::Tensor C, torch::Tensor A, torch::Tensor B,+ torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad);+ '''++ # =============================================================================+ # PIPELINED KERNEL MODULE (exact copy from v14)+ # =============================================================================++ PIPELINED_CUDA = r'''+ #include <torch/extension.h>+ #include <cuda_fp16.h>+ #include <cuda_fp4.h>+ #include <cuda_fp8.h>+#define ASYNC_CP_16(dst, src) \asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" \:: "r"((uint32_t)__cvta_generic_to_shared(dst)), "l"(src))⋯ 5 unchanged lines#define NUM_STAGES 2- __device__ __forceinline__ float dot_scaled_pipelined(+ __device__ __forceinline__ void cvt_fp4_to_f16(half2* out, uint32_t in) {+ asm volatile(+ "{\n\t"+ ".reg .b8 b0, b1, b2, b3;\n\t"+ "mov.b32 {b0, b1, b2, b3}, %4;\n\t"+ "cvt.rn.f16x2.e2m1x2 %0, b0;\n\t"+ "cvt.rn.f16x2.e2m1x2 %1, b1;\n\t"+ "cvt.rn.f16x2.e2m1x2 %2, b2;\n\t"+ "cvt.rn.f16x2.e2m1x2 %3, b3;\n\t"+ "}"+ : "=r"(reinterpret_cast<int&>(out[0])),+ "=r"(reinterpret_cast<int&>(out[1])),+ "=r"(reinterpret_cast<int&>(out[2])),+ "=r"(reinterpret_cast<int&>(out[3]))+ : "r"(in)+ );+ }++ __device__ __forceinline__ half2 cvt_fp8_to_f16(uint16_t in) {+ int out;+ asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(in));+ return *reinterpret_cast<half2*>(&out);+ }++ __device__ __forceinline__ float dot_scaled_32fp4(const uint8_t* a_smem,const uint8_t* b_smem,half2 scale⋯ 39 unchanged lineshalf* __restrict__ c,int M, int K, int L, int N_pad) {- constexpr int BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW;constexpr int K_TILE_BYTES = K_TILE / 2;constexpr int K_TILE_SF = K_TILE / 16;⋯ 86 unchanged linesuint16_t sfb_packed = *reinterpret_cast<const uint16_t*>(&sh_sfb[stage][sf_idx]);half2 scale = __hmul2(cvt_fp8_to_f16(sfa_packed), cvt_fp8_to_f16(sfb_packed));- acc += dot_scaled_pipelined(+ acc += dot_scaled_32fp4(&sh_a[stage][warp_id][byte_idx],&sh_b[stage][byte_idx],scale⋯ 16 unchanged lines}}- // =============================================================================- // Dispatch- // =============================================================================+ template<int ROWS, int THREADS, int K_TILE>+ void dispatch(torch::Tensor& c, torch::Tensor& a, torch::Tensor& b,+ torch::Tensor& sfa, torch::Tensor& sfb, int m, int k, int l, int n) {+ dim3 grid((m + ROWS - 1) / ROWS, l);+ dim3 block(ROWS * THREADS);+ gemv_pipelined_kernel<ROWS, THREADS, K_TILE><<<grid, block>>>(+ a.data_ptr<uint8_t>(), b.data_ptr<uint8_t>(),+ sfa.data_ptr<uint8_t>(), sfb.data_ptr<uint8_t>(),+ reinterpret_cast<half*>(c.data_ptr<at::Half>()), m, k, l, n);+ }- void run_nvfp4_gemv(+ void run_pipelined(torch::Tensor C, torch::Tensor A, torch::Tensor B,torch::Tensor SFA, torch::Tensor SFB,- int m, int k, int l, int n_pad+ int m, int k, int l, int n_pad, int k_tile_hint) {- constexpr auto CS = CacheHint::CS;- constexpr auto CA = CacheHint::CA;-- // Selection heuristic based on benchmarks:- // - k=16384, l=1: pipelined faster (24.6 vs 30.8)- // - k=7168, l=8: direct faster (38.5 vs 51.0)- // - k=2048, l=4: same (16.4)-- const bool use_pipelined = (k >= 8192 && l <= 2) || // Large K, small batch- (k <= 4096 && l >= 4); // Small K, larger batch-- const int k_bytes = k / 2;- const bool b_fits_l1 = (k_bytes <= 32 * 1024);-- auto* a_ptr = A.data_ptr<uint8_t>();- auto* b_ptr = B.data_ptr<uint8_t>();- auto* sfa_ptr = SFA.data_ptr<uint8_t>();- auto* sfb_ptr = SFB.data_ptr<uint8_t>();- auto* c_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());-- if (use_pipelined) {- // Pipelined kernel - configs from v14- if (k >= 8192) {- dim3 grid((m + 4 - 1) / 4, l);- dim3 block(4 * 32);- gemv_pipelined_kernel<4, 32, 2048><<<grid, block>>>(- a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);- } else if (k >= 4096) {- dim3 grid((m + 4 - 1) / 4, l);- dim3 block(4 * 32);- gemv_pipelined_kernel<4, 32, 1024><<<grid, block>>>(- a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);- } else {- dim3 grid((m + 8 - 1) / 8, l);- dim3 block(8 * 16);- gemv_pipelined_kernel<8, 16, 2048><<<grid, block>>>(- a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ if (k >= 8192) {+ // Try different K_TILE values for large K+ // k_tile_hint: 0=2048, 1=1024, 2=4096, 3=8192+ switch (k_tile_hint) {+ case 1: dispatch<4, 32, 1024>(C, A, B, SFA, SFB, m, k, l, n_pad); break;+ case 2: dispatch<4, 32, 4096>(C, A, B, SFA, SFB, m, k, l, n_pad); break;+ case 3: dispatch<4, 32, 8192>(C, A, B, SFA, SFB, m, k, l, n_pad); break;+ default: dispatch<4, 32, 2048>(C, A, B, SFA, SFB, m, k, l, n_pad); break;}+ } else if (k >= 4096) {+ dispatch<4, 32, 1024>(C, A, B, SFA, SFB, m, k, l, n_pad);} else {- // Direct kernel - configs from v13- if (k >= 8192) {- dim3 grid((m + 8 - 1) / 8, l);- dim3 block(8 * 32);- if (b_fits_l1) {- gemv_direct_kernel<8, 32, CS, CA><<<grid, block>>>(- a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);- } else {- gemv_direct_kernel<8, 32, CS, CS><<<grid, block>>>(- a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);- }- } else if (k >= 4096) {- dim3 grid((m + 4 - 1) / 4, l);- dim3 block(4 * 32);- gemv_direct_kernel<4, 32, CS, CA><<<grid, block>>>(- a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);- } else {- dim3 grid((m + 8 - 1) / 8, l);- dim3 block(8 * 16);- gemv_direct_kernel<8, 16, CS, CA><<<grid, block>>>(- a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);- }+ dispatch<8, 16, 2048>(C, A, B, SFA, SFB, m, k, l, n_pad);}}'''- CPP_SRC = r'''+ PIPELINED_CPP = r'''#include <torch/extension.h>- void run_nvfp4_gemv(torch::Tensor C, torch::Tensor A, torch::Tensor B,- torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad);+ void run_pipelined(torch::Tensor C, torch::Tensor A, torch::Tensor B,+ torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad, int k_tile_hint);'''- _module = None+ # =============================================================================+ # Module loading+ # =============================================================================- def _load():- global _module- if _module is None:- _module = load_inline(- name='nvfp4_gemv_v15_hybrid',- cpp_sources=CPP_SRC,- cuda_sources=CUDA_SRC,- functions=['run_nvfp4_gemv'],+ # =============================================================================+ # Configuration+ # =============================================================================+ # K_TILE options for k >= 8192:+ # 0 = 2048 (default)+ # 1 = 1024+ # 2 = 4096+ # 3 = 8192+ K_TILE_HINT = 0 # <-- Change this to test different K_TILE values++ _direct_module = None+ _pipelined_module = None++ def _load_direct():+ global _direct_module+ if _direct_module is None:+ _direct_module = load_inline(+ name='nvfp4_gemv_direct_v16',+ cpp_sources=DIRECT_CPP,+ cuda_sources=DIRECT_CUDA,+ functions=['run_direct'],extra_cuda_cflags=['-O3', '--use_fast_math', '-std=c++17','-gencode=arch=compute_100a,code=sm_100a',],verbose=False,)- return _module+ return _direct_module+ def _load_pipelined():+ global _pipelined_module+ if _pipelined_module is None:+ _pipelined_module = load_inline(+ name='nvfp4_gemv_pipelined_v16b', # Updated to force recompile+ cpp_sources=PIPELINED_CPP,+ cuda_sources=PIPELINED_CUDA,+ functions=['run_pipelined'],+ extra_cuda_cflags=[+ '-O3', '--use_fast_math', '-std=c++17',+ '-gencode=arch=compute_100a,code=sm_100a',+ ],+ verbose=False,+ )+ return _pipelined_module+def custom_kernel(data: input_t) -> output_t:a, b, sfa, sfb, _, _, c = datam, k_packed, l = a.shape⋯ 5 unchanged linesif not sfb.is_cuda:sfb = sfb.to(a.device)- _load().run_nvfp4_gemv(- c.squeeze(1), a.view(torch.uint8), b.view(torch.uint8),- sfa.view(torch.uint8), sfb.view(torch.uint8), m, k, l, n_pad)+ c_out = c.squeeze(1)+ a_u8 = a.view(torch.uint8)+ b_u8 = b.view(torch.uint8)+ sfa_u8 = sfa.view(torch.uint8)+ sfb_u8 = sfb.view(torch.uint8)++ # Selection heuristic based on benchmarks:+ # - k=16384, l=1: pipelined faster (24.6 vs 30.8)+ # - k=7168, l=8: direct faster (38.5 vs 51.0)+ # - k=2048, l=4: same (16.4)++ use_pipelined = (k >= 8192 and l <= 2) or (k <= 4096 and l >= 4)++ if use_pipelined:+ _load_pipelined().run_pipelined(c_out, a_u8, b_u8, sfa_u8, sfb_u8, m, k, l, n_pad, 2)+ else:+ _load_direct().run_direct(c_out, a_u8, b_u8, sfa_u8, sfb_u8, m, k, l, n_pad)+return c+
scrolls · 369 diff lines total
Best evidence level for this revision: reported
JSON