submission 109980
macto · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 452 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-109980?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:6b570c84858031bba9935e2cb012f78c345de7f1edda0390da304074bb46e4c3
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.py452 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# Combine both v13 (direct) and v14 (pipelined) kernels with selection logic
CUDA_SRC = 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>
__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);
}
// =============================================================================
// DIRECT KERNEL (from v13)
// =============================================================================
__device__ __forceinline__ float dot_scaled_direct(
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_direct_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_direct(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);
}
}
// =============================================================================
// PIPELINED KERNEL (from v14)
// =============================================================================
#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__ float dot_scaled_pipelined(
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 BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW;
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_pipelined(
&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);
}
}
// =============================================================================
// Dispatch
// =============================================================================
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
) {
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);
}
} 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);
}
}
}
'''
CPP_SRC = 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);
'''
_module = None
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'],
extra_cuda_cflags=[
'-O3', '--use_fast_math', '-std=c++17',
'-gencode=arch=compute_100a,code=sm_100a',
],
verbose=False,
)
return _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)
_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)
return c
scrolls · 452 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 109911.
⋯ 1 unchanged linesfrom torch.utils.cpp_extension import load_inlinefrom task import input_t, output_t+ # Combine both v13 (direct) and v14 (pipelined) kernels with selection logicCUDA_SRC = r'''#include <torch/extension.h>#include <cuda_fp16.h>#include <cuda_fp4.h>#include <cuda_fp8.h>+ // =============================================================================+ // Common Utilities+ // =============================================================================- __device__ __forceinline__ void ldcs_u32x4(uint32_t* dst, const void* src) {- 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));- }+ enum class CacheHint { CS, CA };- // Cached load for B (keep in L1, reused across M rows)- __device__ __forceinline__ void ldca_u32x4(uint32_t* dst, const void* src) {- 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_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));+ }}- __device__ __forceinline__ void ldcs_u16(uint16_t* dst, const void* src) {- asm volatile("ld.global.cs.u16 %0, [%1];" : "=h"(*dst) : "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));+ }}- // Cached load for scale factors B- __device__ __forceinline__ void ldca_u16(uint16_t* dst, const void* src) {- asm volatile("ld.global.ca.u16 %0, [%1];" : "=h"(*dst) : "l"(src));- }--- // Convert 4 bytes of FP4x2 to 4 x half2 (8 FP16 values)- __device__ __forceinline__ void fp4x8_to_half2x4(half2* out, uint32_t in) {+ __device__ __forceinline__ void cvt_fp4_to_f16(half2* out, uint32_t in) {asm volatile("{\n\t"".reg .b8 b0, b1, b2, b3;\n\t"⋯ 11 unchanged lines);}- // Convert 2 FP8 scale factors to half2- __device__ __forceinline__ half2 fp8x2_to_half2(uint16_t 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 blockscaled_dot_16bytes(- const uint32_t* a_regs, // 4 x uint32 = 16 bytes = 32 FP4- const uint32_t* b_regs, // 4 x uint32 = 16 bytes = 32 FP4- half2 combined_scale // SFA * SFB already fused+ // =============================================================================+ // DIRECT KERNEL (from v13)+ // =============================================================================++ __device__ __forceinline__ float dot_scaled_direct(+ const uint32_t* a_data,+ const uint32_t* b_data,+ half2 scale) {- // Unpack to half2- half2 a_h2[16], b_h2[16];- fp4x8_to_half2x4(a_h2 + 0, a_regs[0]);- fp4x8_to_half2x4(a_h2 + 4, a_regs[1]);- fp4x8_to_half2x4(a_h2 + 8, a_regs[2]);- fp4x8_to_half2x4(a_h2 + 12, a_regs[3]);+ 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]);+ }- fp4x8_to_half2x4(b_h2 + 0, b_regs[0]);- fp4x8_to_half2x4(b_h2 + 4, b_regs[1]);- fp4x8_to_half2x4(b_h2 + 8, b_regs[2]);- fp4x8_to_half2x4(b_h2 + 12, b_regs[3]);+ half2 acc0 = __hmul2(a_f16[0], b_f16[0]);+ half2 acc1 = __hmul2(a_f16[8], b_f16[8]);- // Dot product WITHOUT scale (deferred scaling)- // First 8 half2 use scale.x, next 8 use scale.y- half2 acc0 = __hmul2(a_h2[0], b_h2[0]);- half2 acc1 = __hmul2(a_h2[8], b_h2[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_direct_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_direct(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);+ }+ }++ // =============================================================================+ // PIPELINED KERNEL (from v14)+ // =============================================================================++ #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__ float dot_scaled_pipelined(+ 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 unrollfor (int i = 1; i < 8; i++) {- acc0 = __hfma2(a_h2[i], b_h2[i], acc0);- acc1 = __hfma2(a_h2[8 + i], b_h2[8 + i], acc1);+ acc0 = __hfma2(a_f16[i], b_f16[i], acc0);+ acc1 = __hfma2(a_f16[8 + i], b_f16[8 + i], acc1);}- // Reduce each half2 to halfhalf 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"(*reinterpret_cast<uint16_t*>(&sum0)), "h"(*reinterpret_cast<uint16_t*>(&combined_scale.x)));- asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;" : "+f"(result) : "h"(*reinterpret_cast<uint16_t*>(&sum1)), "h"(*reinterpret_cast<uint16_t*>(&combined_scale.y)));+ 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 BLOCK_M, int THREADS_PER_ROW>- __global__ __launch_bounds__(BLOCK_M * THREADS_PER_ROW)- void gemv_optimized_kernel(+ 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_rows+ int M, int K, int L, int N_pad) {- constexpr int BLOCK_SIZE = BLOCK_M * THREADS_PER_ROW;+ 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;- const int m_base = blockIdx.x * BLOCK_M;- const int batch_id = blockIdx.y;+ __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 m = m_base + warp_id;- if (m >= M) return;+ 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_batch_stride = (size_t)M * K_bytes;- const size_t b_batch_stride = (size_t)N_rows * K_bytes;- const size_t sfa_batch_stride = (size_t)M * K_sf;- const size_t sfb_batch_stride = (size_t)N_rows * K_sf;+ 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;- const uint8_t* row_a = a + batch_id * a_batch_stride + m * K_bytes;- const uint8_t* batch_b = b + batch_id * b_batch_stride;- const uint8_t* row_sfa = sfa + batch_id * sfa_batch_stride + m * K_sf;- const uint8_t* batch_sfb = sfb + batch_id * sfb_batch_stride;+ 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;- // Process 2 scale groups per iteration (32 FP4 values = 16 bytes)- // Each thread handles multiple scale pairs- const int num_scale_pairs = K_sf / 2;-- #pragma unroll 4- for (int sp = lane; sp < num_scale_pairs; sp += THREADS_PER_ROW) {- int sf_base = sp * 2;- int byte_base = sf_base * 8; // 8 bytes per scale factor+ for (int tile = 0; tile < num_tiles; tile++) {+ const int stage = tile % NUM_STAGES;+ const int next_stage = (tile + 1) % NUM_STAGES;- // Load scale factors with cache hints- uint16_t sfa_raw, sfb_raw;- ldcs_u16(&sfa_raw, row_sfa + sf_base);- ldca_u16(&sfb_raw, batch_sfb + sf_base);+ if (tile + 1 < num_tiles) {+ issue_tile(tile + 1, next_stage);+ }- // Convert FP8x2 to half2 and fuse scales- half2 scale_a = fp8x2_to_half2(sfa_raw);- half2 scale_b = fp8x2_to_half2(sfb_raw);- half2 combined_scale = __hmul2(scale_a, scale_b);+ 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;- // Load 16 bytes of A and B with cache hints- uint32_t a_regs[4], b_regs[4];- ldcs_u32x4(a_regs, row_a + byte_base);- ldca_u32x4(b_regs, batch_b + byte_base);+ #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_pipelined(+ &sh_a[stage][warp_id][byte_idx],+ &sh_b[stage][byte_idx],+ scale+ );+ }- // Compute with deferred scaling- acc += blockscaled_dot_16bytes(a_regs, b_regs, combined_scale);+ if (tile + 1 < num_tiles) {+ ASYNC_WAIT(0);+ __syncthreads();+ }}- // Warp reduction#pragma unrollfor (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {acc += __shfl_down_sync(0xffffffff, acc, offset);}if (lane == 0) {- c[(size_t)m + (size_t)batch_id * M] = __float2half(acc);+ c[(size_t)row + (size_t)batch * M] = __float2half(acc);}}+ // =============================================================================+ // Dispatch+ // =============================================================================+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) {- // K-specialized dispatch- if (k >= 8192) {- // Large K: 8 rows per block, 32 threads per row- constexpr int BLOCK_M = 8;- constexpr int THREADS_PER_ROW = 32;- int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;- dim3 grid(num_m_blocks, l);- dim3 block(BLOCK_M * THREADS_PER_ROW);- gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<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_pad);- } else if (k >= 4096) {- // Medium-large K: 4 rows per block, 32 threads per row- constexpr int BLOCK_M = 4;- constexpr int THREADS_PER_ROW = 32;- int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;- dim3 grid(num_m_blocks, l);- dim3 block(BLOCK_M * THREADS_PER_ROW);- gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<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_pad);+ 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);+ }} else {- // Small K: 8 rows per block, 16 threads per row- constexpr int BLOCK_M = 8;- constexpr int THREADS_PER_ROW = 16;- int num_m_blocks = (m + BLOCK_M - 1) / BLOCK_M;- dim3 grid(num_m_blocks, l);- dim3 block(BLOCK_M * THREADS_PER_ROW);- gemv_optimized_kernel<BLOCK_M, THREADS_PER_ROW><<<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_pad);+ // 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);+ }}}'''⋯ 4 unchanged linestorch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad);'''- _cuda_module = None+ _module = None- def get_cuda_module():- global _cuda_module- if _cuda_module is None:- _cuda_module = load_inline(- name='nvfp4_gemv_optimized_v13',+ 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'],extra_cuda_cflags=[- '-O3',- '--use_fast_math',- '-std=c++17',+ '-O3', '--use_fast_math', '-std=c++17','-gencode=arch=compute_100a,code=sm_100a',],verbose=False,)- return _cuda_module+ return _moduledef custom_kernel(data: input_t) -> output_t:- a, b, sfa_ref, sfb_ref, _, _, c = data- module = get_cuda_module()+ a, b, sfa, sfb, _, _, c = datam, k_packed, l = a.shapek = k_packed * 2n_pad = b.shape[0]- if not sfa_ref.is_cuda:- sfa_ref = sfa_ref.to(a.device)- if not sfb_ref.is_cuda:- sfb_ref = sfb_ref.to(a.device)- c_out = c.squeeze(1)- module.run_nvfp4_gemv(c_out, a.view(torch.uint8), b.view(torch.uint8),- sfa_ref.view(torch.uint8), sfb_ref.view(torch.uint8),- m, k, l, n_pad)++ if not sfa.is_cuda:+ sfa = sfa.to(a.device)+ if 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)return c-
scrolls · 581 diff lines total
Best evidence level for this revision: reported
JSON