submission 106932
macto · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 448 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-106932?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:302a2b750e37154d2d5ad64a3ad598535c9886c39f438a968ef267ad7da347bb
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;" :: "r"(dst), "l"(src))fp8
__nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);shared-memory
__shared__ __align__(16) uint8_t sh_a[NUM_BUFFERS][BYTES_PER_TILE];vector-width = float2
float2 f0 = __half22float2(acc_h2_0);Kernel source
submission.py448 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CUDA_SRC = r'''
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#define BLOCK_SIZE 32
#define K_TILE 2560 // Larger tile = fewer iterations
#define SCALES_PER_TILE (K_TILE / 16) // 160
#define BYTES_PER_TILE (K_TILE / 2) // 1280
#define NUM_BUFFERS 2
__device__ __forceinline__ __half2 decode_fp4x2(uint8_t byte) {
__half2_raw raw = __nv_cvt_fp4x2_to_halfraw2(
static_cast<__nv_fp4x2_storage_t>(byte),
__NV_E2M1
);
return *reinterpret_cast<__half2*>(&raw);
}
__device__ __forceinline__ float decode_fp8(int8_t byte) {
__nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);
__half_raw raw = __nv_cvt_fp8_to_halfraw(storage, __NV_E4M3);
return __half2float(__ushort_as_half(raw.x));
}
__device__ __forceinline__ __half2 dot_scaled_4bytes(
uint32_t a4,
uint32_t b4,
__half2 scale_h2
) {
uint32_t b_byte0, b_byte1, b_byte2, b_byte3;
uint32_t a_byte0, a_byte1, a_byte2, a_byte3;
asm("bfe.u32 %0, %1, 0, 8;" : "=r"(b_byte0) : "r"(b4));
asm("bfe.u32 %0, %1, 8, 8;" : "=r"(b_byte1) : "r"(b4));
asm("bfe.u32 %0, %1, 16, 8;" : "=r"(b_byte2) : "r"(b4));
asm("bfe.u32 %0, %1, 24, 8;" : "=r"(b_byte3) : "r"(b4));
asm("bfe.u32 %0, %1, 0, 8;" : "=r"(a_byte0) : "r"(a4));
asm("bfe.u32 %0, %1, 8, 8;" : "=r"(a_byte1) : "r"(a4));
asm("bfe.u32 %0, %1, 16, 8;" : "=r"(a_byte2) : "r"(a4));
asm("bfe.u32 %0, %1, 24, 8;" : "=r"(a_byte3) : "r"(a4));
__half2 acc = __hmul2(decode_fp4x2(a_byte0), __hmul2(decode_fp4x2(b_byte0), scale_h2));
acc = __hfma2(decode_fp4x2(a_byte1), __hmul2(decode_fp4x2(b_byte1), scale_h2), acc);
acc = __hfma2(decode_fp4x2(a_byte2), __hmul2(decode_fp4x2(b_byte2), scale_h2), acc);
acc = __hfma2(decode_fp4x2(a_byte3), __hmul2(decode_fp4x2(b_byte3), scale_h2), acc);
return acc;
}
__device__ __forceinline__ float compute_tile(
const uint8_t* sh_a,
const uint8_t* sh_b,
const uint8_t* sh_sfa,
const uint8_t* sh_sfb,
int tid
) {
float acc = 0.0f;
#pragma unroll 8
for (int sf = tid; sf < SCALES_PER_TILE; sf += BLOCK_SIZE) {
float scale = decode_fp8(static_cast<int8_t>(sh_sfa[sf])) *
decode_fp8(static_cast<int8_t>(sh_sfb[sf]));
__half2 scale_h2 = __half2half2(__float2half(scale));
int byte_base = sf << 3; // sf * 8
uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base]);
uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base]);
uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base + 4]);
uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base + 4]);
__half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
__half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);
float2 f0 = __half22float2(acc_h2_0);
float2 f1 = __half22float2(acc_h2_1);
acc += f0.x + f0.y + f1.x + f1.y;
}
return acc;
}
// ============================================================================
// Compute remainder from global memory
// ============================================================================
__device__ __forceinline__ float compute_remainder(
const uint8_t* row_a,
const uint8_t* batch_b,
const uint8_t* row_sfa,
const uint8_t* batch_sfb,
int remainder_sf_start,
int K_sf,
int tid
) {
float acc = 0.0f;
#pragma unroll 4
for (int sf = remainder_sf_start + tid; sf < K_sf; sf += BLOCK_SIZE) {
float scale = decode_fp8(static_cast<int8_t>(__ldg(&row_sfa[sf]))) *
decode_fp8(static_cast<int8_t>(__ldg(&batch_sfb[sf])));
__half2 scale_h2 = __half2half2(__float2half(scale));
int byte_base = sf << 3;
uint32_t a4_0 = __ldg(reinterpret_cast<const uint32_t*>(&row_a[byte_base]));
uint32_t b4_0 = __ldg(reinterpret_cast<const uint32_t*>(&batch_b[byte_base]));
uint32_t a4_1 = __ldg(reinterpret_cast<const uint32_t*>(&row_a[byte_base + 4]));
uint32_t b4_1 = __ldg(reinterpret_cast<const uint32_t*>(&batch_b[byte_base + 4]));
__half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
__half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);
float2 f0 = __half22float2(acc_h2_0);
float2 f1 = __half22float2(acc_h2_1);
acc += f0.x + f0.y + f1.x + f1.y;
}
return acc;
}
__device__ __forceinline__ float compute_remainder_smem(
const uint8_t* sh_a,
const uint8_t* sh_b,
const uint8_t* sh_sfa,
const uint8_t* sh_sfb,
int remainder_scales,
int tid
) {
float acc = 0.0f;
#pragma unroll 4
for (int sf = tid; sf < remainder_scales; sf += BLOCK_SIZE) {
float scale = decode_fp8(static_cast<int8_t>(sh_sfa[sf])) *
decode_fp8(static_cast<int8_t>(sh_sfb[sf]));
__half2 scale_h2 = __half2half2(__float2half(scale));
int byte_base = sf << 3;
uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base]);
uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base]);
uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base + 4]);
uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base + 4]);
__half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);
__half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);
float2 f0 = __half22float2(acc_h2_0);
float2 f1 = __half22float2(acc_h2_1);
acc += f0.x + f0.y + f1.x + f1.y;
}
return acc;
}
// ============================================================================
// Main GEMV kernel - one block per (row, batch)
// ============================================================================
__global__ __launch_bounds__(BLOCK_SIZE)
void gemv_nvfp4_kernel(
const int8_t* __restrict__ a,
const int8_t* __restrict__ b,
const int8_t* __restrict__ sfa,
const int8_t* __restrict__ sfb,
half* __restrict__ c,
int M, int K, int L,
int N_rows
) {
// Double-buffered shared memory
__shared__ __align__(16) uint8_t sh_a[NUM_BUFFERS][BYTES_PER_TILE];
__shared__ __align__(16) uint8_t sh_b[NUM_BUFFERS][BYTES_PER_TILE];
__shared__ __align__(16) uint8_t sh_sfa[NUM_BUFFERS][SCALES_PER_TILE];
__shared__ __align__(16) uint8_t sh_sfb[NUM_BUFFERS][SCALES_PER_TILE];
__shared__ float smem_acc[BLOCK_SIZE / 32];
const int m = blockIdx.x;
const int l = blockIdx.y;
const int tid = threadIdx.x;
if (m >= M) return;
// Dimension calculations
const int K_bytes = K / 2;
const int K_sf = K / 16;
// Batch strides
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;
// Pointers for this row and batch
const uint8_t* row_a = reinterpret_cast<const uint8_t*>(a) + l * a_batch_stride + m * K_bytes;
const uint8_t* batch_b = reinterpret_cast<const uint8_t*>(b) + l * b_batch_stride;
const uint8_t* row_sfa = reinterpret_cast<const uint8_t*>(sfa) + l * sfa_batch_stride + m * K_sf;
const uint8_t* batch_sfb = reinterpret_cast<const uint8_t*>(sfb) + l * sfb_batch_stride;
// Tile counts
const int tile_count = K_bytes / BYTES_PER_TILE;
const int remainder_start = tile_count * BYTES_PER_TILE;
float acc = 0.0f;
// Async copy macros
#define ASYNC_COPY_16(dst, src) \
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))
#define ASYNC_COPY_4(dst, src) \
asm volatile("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"(dst), "l"(src))
// Lambda for issuing tile async copies
auto issue_tile_async = [&](int b_idx, int tile) {
const int base_byte = tile * BYTES_PER_TILE;
const int base_sf = tile * SCALES_PER_TILE;
const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][0]);
const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][0]);
const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);
// Copy A and B tiles (BYTES_PER_TILE bytes each)
for (int i = tid * 16; i < BYTES_PER_TILE; i += BLOCK_SIZE * 16) {
ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
}
// Copy scale factors (SCALES_PER_TILE bytes each)
for (int i = tid * 4; i < SCALES_PER_TILE; i += BLOCK_SIZE * 4) {
ASYNC_COPY_4(sh_sfa_base + i, row_sfa + base_sf + i);
ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + base_sf + i);
}
asm volatile("cp.async.commit_group;");
};
auto issue_remainder_async = [&](int b_idx, int rem_sf_start, int total_K_sf) {
const int base_byte = rem_sf_start << 3; // rem_sf_start * 8
const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][0]);
const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);
const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][0]);
const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);
int remainder_bytes = (total_K_sf - rem_sf_start) << 3; // * 8
int remainder_scales = total_K_sf - rem_sf_start;
for (int i = tid * 16; i < remainder_bytes; i += BLOCK_SIZE * 16) {
ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);
ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);
}
for (int i = tid * 4; i < remainder_scales; i += BLOCK_SIZE * 4) {
ASYNC_COPY_4(sh_sfa_base + i, row_sfa + rem_sf_start + i);
ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + rem_sf_start + i);
}
asm volatile("cp.async.commit_group;");
};
// Main loop
int remainder_sf_start = remainder_start / 8; // Convert bytes to scale factor index
bool has_remainder = (remainder_sf_start < K_sf);
int buf = 0;
if (tile_count > 0) {
// Load first tile
issue_tile_async(0, 0);
asm volatile("cp.async.wait_group 0;");
__syncthreads();
for (int tile = 0; tile < tile_count; ++tile) {
if (tile + 1 < tile_count) {
// Prefetch next tile
issue_tile_async(buf ^ 1, tile + 1);
} else if (has_remainder) {
// On last tile: prefetch remainder
issue_remainder_async(buf ^ 1, remainder_sf_start, K_sf);
}
// Compute current tile
acc += compute_tile(
sh_a[buf],
sh_b[buf],
sh_sfa[buf],
sh_sfb[buf],
tid
);
if (tile + 1 < tile_count || has_remainder) {
asm volatile("cp.async.wait_group 0;");
__syncthreads();
buf ^= 1;
}
}
}
if (has_remainder) {
if (tile_count > 0) {
int remainder_scales = K_sf - remainder_sf_start;
acc += compute_remainder_smem(
sh_a[buf],
sh_b[buf],
sh_sfa[buf],
sh_sfb[buf],
remainder_scales,
tid
);
} else {
// No tiles - load remainder directly from global memory
acc += compute_remainder(
row_a,
batch_b,
row_sfa,
batch_sfb,
remainder_sf_start,
K_sf,
tid
);
}
}
#undef ASYNC_COPY_16
#undef ASYNC_COPY_4
// ========================================================================
// Warp-level reduction
// ========================================================================
float warp_sum = acc;
warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 16);
warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 8);
warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 4);
warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 2);
warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 1);
// ========================================================================
// Block-level reduction (only 1 warp, so just write)
// ========================================================================
const int warp_id = tid >> 5;
const int lane = tid & 31;
if (lane == 0) {
smem_acc[warp_id] = warp_sum;
}
__syncthreads();
if (warp_id == 0) {
float block_sum = (lane < (BLOCK_SIZE >> 5)) ? smem_acc[lane] : 0.0f;
block_sum += __shfl_down_sync(0xffffffff, block_sum, 16);
block_sum += __shfl_down_sync(0xffffffff, block_sum, 8);
block_sum += __shfl_down_sync(0xffffffff, block_sum, 4);
block_sum += __shfl_down_sync(0xffffffff, block_sum, 2);
block_sum += __shfl_down_sync(0xffffffff, block_sum, 1);
// ====================================================================
// Final output write
// ====================================================================
if (lane == 0) {
size_t c_idx = (size_t)m + (size_t)l * M;
c[c_idx] = __float2half(block_sum);
}
}
}
// ============================================================================
// Host wrapper
// ============================================================================
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
) {
dim3 grid(m, l);
dim3 block(BLOCK_SIZE);
gemv_nvfp4_kernel<<<grid, block>>>(
reinterpret_cast<const int8_t*>(A.data_ptr()),
reinterpret_cast<const int8_t*>(B.data_ptr()),
reinterpret_cast<const int8_t*>(SFA.data_ptr()),
reinterpret_cast<const int8_t*>(SFB.data_ptr()),
reinterpret_cast<half*>(C.data_ptr<at::Half>()),
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
);
'''
_cuda_module = None
def get_cuda_module():
global _cuda_module
if _cuda_module is None:
_cuda_module = load_inline(
name='nvfp4_gemv_async_v2',
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 _cuda_module
def custom_kernel(data: input_t) -> output_t:
a, b, sfa_ref, sfb_ref, _, _, c = data
module = get_cuda_module()
m, k_packed, l = a.shape
k = k_packed * 2
n_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)
a_uint8 = a.view(torch.uint8)
b_uint8 = b.view(torch.uint8)
sfa_uint8 = sfa_ref.view(torch.uint8)
sfb_uint8 = sfb_ref.view(torch.uint8)
c_out = c.squeeze(1)
module.run_nvfp4_gemv(
c_out, a_uint8, b_uint8, sfa_uint8, sfb_uint8,
m, k, l, n_pad
)
return c
scrolls · 448 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 105811.
- """- Enhanced GEMV based on our working baseline with selective improvements:- 1. Native FP4/FP8 conversion functions- 2. Double buffering for B and SFB with async copy- 3. Keep our thread block structure (32x32 = 1024 threads)- 4. Keep our per-thread work (16 bytes per iteration)- 5. SM100a architecture flag for Blackwell- """-import torchfrom torch.utils.cpp_extension import load_inlinefrom task import input_t, output_tCUDA_SRC = r'''#include <torch/extension.h>- #include <cuda_runtime.h>#include <cuda_fp16.h>#include <cuda_fp4.h>#include <cuda_fp8.h>- #include <cstdint>- #define THREADS_PER_M 32- #define THREADS_PER_K 32- #define BLOCK_SIZE (THREADS_PER_M * THREADS_PER_K)- #define SF_VEC_SIZE 16- #define K_TILE_SIZE 512 // 32 threads × 16 bytes- #define SF_TILE_SIZE 64 // K_TILE_SIZE / 8+ #define BLOCK_SIZE 32+ #define K_TILE 2560 // Larger tile = fewer iterations+ #define SCALES_PER_TILE (K_TILE / 16) // 160+ #define BYTES_PER_TILE (K_TILE / 2) // 1280+ #define NUM_BUFFERS 2- // Native FP4 E2M1 to half2 conversion (from reference)__device__ __forceinline__ __half2 decode_fp4x2(uint8_t byte) {__half2_raw raw = __nv_cvt_fp4x2_to_halfraw2(static_cast<__nv_fp4x2_storage_t>(byte),⋯ 2 unchanged linesreturn *reinterpret_cast<__half2*>(&raw);}- // Native FP8 E4M3 to float conversion (from reference)- __device__ __forceinline__ float decode_fp8(uint8_t byte) {+ __device__ __forceinline__ float decode_fp8(int8_t byte) {__nv_fp8_storage_t storage = static_cast<__nv_fp8_storage_t>(byte);__half_raw raw = __nv_cvt_fp8_to_halfraw(storage, __NV_E4M3);return __half2float(__ushort_as_half(raw.x));}+ __device__ __forceinline__ __half2 dot_scaled_4bytes(+ uint32_t a4,+ uint32_t b4,+ __half2 scale_h2+ ) {+ uint32_t b_byte0, b_byte1, b_byte2, b_byte3;+ uint32_t a_byte0, a_byte1, a_byte2, a_byte3;- // Async copy macros- #define ASYNC_COPY_16(dst, src) \- asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))+ asm("bfe.u32 %0, %1, 0, 8;" : "=r"(b_byte0) : "r"(b4));+ asm("bfe.u32 %0, %1, 8, 8;" : "=r"(b_byte1) : "r"(b4));+ asm("bfe.u32 %0, %1, 16, 8;" : "=r"(b_byte2) : "r"(b4));+ asm("bfe.u32 %0, %1, 24, 8;" : "=r"(b_byte3) : "r"(b4));- #define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")- #define ASYNC_WAIT_ALL() asm volatile("cp.async.wait_group 0;")+ asm("bfe.u32 %0, %1, 0, 8;" : "=r"(a_byte0) : "r"(a4));+ asm("bfe.u32 %0, %1, 8, 8;" : "=r"(a_byte1) : "r"(a4));+ asm("bfe.u32 %0, %1, 16, 8;" : "=r"(a_byte2) : "r"(a4));+ asm("bfe.u32 %0, %1, 24, 8;" : "=r"(a_byte3) : "r"(a4));- // Kernel for l == 1 with double-buffered B- __global__ __launch_bounds__(BLOCK_SIZE, 2)- void nvfp4_gemv_kernel_single(- half* __restrict__ C,- const uint8_t* __restrict__ A,- const uint8_t* __restrict__ B,- const uint8_t* __restrict__ SFA,- const uint8_t* __restrict__ SFB,- int m, int k_packed, int k_sf, int n_pad+ __half2 acc = __hmul2(decode_fp4x2(a_byte0), __hmul2(decode_fp4x2(b_byte0), scale_h2));+ acc = __hfma2(decode_fp4x2(a_byte1), __hmul2(decode_fp4x2(b_byte1), scale_h2), acc);+ acc = __hfma2(decode_fp4x2(a_byte2), __hmul2(decode_fp4x2(b_byte2), scale_h2), acc);+ acc = __hfma2(decode_fp4x2(a_byte3), __hmul2(decode_fp4x2(b_byte3), scale_h2), acc);++ return acc;+ }++ __device__ __forceinline__ float compute_tile(+ const uint8_t* sh_a,+ const uint8_t* sh_b,+ const uint8_t* sh_sfa,+ const uint8_t* sh_sfb,+ int tid) {- // Double buffer for B (shared across all rows)- __shared__ __align__(16) uint8_t smem_B[2][K_TILE_SIZE];- __shared__ __align__(16) uint8_t smem_SFB[2][SF_TILE_SIZE];-- const int tidx = threadIdx.x;- const int tidy = threadIdx.y;- const int tid = tidy * THREADS_PER_K + tidx;- const int row = blockIdx.x * THREADS_PER_M + tidy;- const bool valid = (row < m);-- const int num_tiles = (k_packed + K_TILE_SIZE - 1) / K_TILE_SIZE;-- float sum = 0.0f;-- // Load first tile- if (num_tiles > 0) {- if (tid < 32) {- int k_off = tid * 16;- if (k_off < k_packed) {- uint32_t dst = __cvta_generic_to_shared(&smem_B[0][tid * 16]);- ASYNC_COPY_16(dst, B + k_off);- }- }- if (tid < 4) {- int sf_off = tid * 16;- if (sf_off < k_sf) {- uint32_t dst = __cvta_generic_to_shared(&smem_SFB[0][tid * 16]);- ASYNC_COPY_16(dst, SFB + sf_off);- }- }- ASYNC_COMMIT();+ float acc = 0.0f;++ #pragma unroll 8+ for (int sf = tid; sf < SCALES_PER_TILE; sf += BLOCK_SIZE) {+ float scale = decode_fp8(static_cast<int8_t>(sh_sfa[sf])) *+ decode_fp8(static_cast<int8_t>(sh_sfb[sf]));+ __half2 scale_h2 = __half2half2(__float2half(scale));++ int byte_base = sf << 3; // sf * 8++ uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base]);+ uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base]);+ uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base + 4]);+ uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base + 4]);++ __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);+ __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);++ float2 f0 = __half22float2(acc_h2_0);+ float2 f1 = __half22float2(acc_h2_1);+ acc += f0.x + f0.y + f1.x + f1.y;}-- for (int tile = 0; tile < num_tiles; tile++) {- int curr = tile & 1;- int next = (tile + 1) & 1;- int next_tile = tile + 1;-- // Prefetch next tile- if (next_tile < num_tiles) {- int k_base_next = next_tile * K_TILE_SIZE;- if (tid < 32) {- int k_off = k_base_next + tid * 16;- if (k_off < k_packed) {- uint32_t dst = __cvta_generic_to_shared(&smem_B[next][tid * 16]);- ASYNC_COPY_16(dst, B + k_off);- }- }- if (tid < 4) {- int sf_off = k_base_next / 8 + tid * 16;- if (sf_off < k_sf) {- uint32_t dst = __cvta_generic_to_shared(&smem_SFB[next][tid * 16]);- ASYNC_COPY_16(dst, SFB + sf_off);- }- }- ASYNC_COMMIT();- }-- // Wait for current tile- ASYNC_WAIT_ALL();- __syncthreads();-- if (valid) {- int k_base = tile * K_TILE_SIZE;- int k_local = tidx * 16;- int k_global = k_base + k_local;-- if (k_global + 16 <= k_packed) {- // Load A directly (per-row, can't share)- uint4 a_vec = *reinterpret_cast<const uint4*>(A + row * k_packed + k_global);- // Load B from shared memory- uint4 b_vec = *reinterpret_cast<const uint4*>(&smem_B[curr][k_local]);-- const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);- const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b_vec);-- int sf_local = k_local / 8;- int sf_global = k_global / 8;-- float sfa0 = decode_fp8(SFA[row * k_sf + sf_global]);- float sfa1 = decode_fp8(SFA[row * k_sf + sf_global + 1]);- float sfb0 = decode_fp8(smem_SFB[curr][sf_local]);- float sfb1 = decode_fp8(smem_SFB[curr][sf_local + 1]);- float scale0 = sfa0 * sfb0;- float scale1 = sfa1 * sfb1;-- __half2 acc0 = __float2half2_rn(0.0f);- __half2 acc1 = __float2half2_rn(0.0f);-- #pragma unroll- for (int i = 0; i < 8; i++) {- acc0 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc0);- }- sum += (__half2float(__low2half(acc0)) + __half2float(__high2half(acc0))) * scale0;-- #pragma unroll- for (int i = 8; i < 16; i++) {- acc1 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc1);- }- sum += (__half2float(__low2half(acc1)) + __half2float(__high2half(acc1))) * scale1;- }- }-- __syncthreads();++ return acc;+ }++ // ============================================================================+ // Compute remainder from global memory+ // ============================================================================+ __device__ __forceinline__ float compute_remainder(+ const uint8_t* row_a,+ const uint8_t* batch_b,+ const uint8_t* row_sfa,+ const uint8_t* batch_sfb,+ int remainder_sf_start,+ int K_sf,+ int tid+ ) {+ float acc = 0.0f;++ #pragma unroll 4+ for (int sf = remainder_sf_start + tid; sf < K_sf; sf += BLOCK_SIZE) {+ float scale = decode_fp8(static_cast<int8_t>(__ldg(&row_sfa[sf]))) *+ decode_fp8(static_cast<int8_t>(__ldg(&batch_sfb[sf])));+ __half2 scale_h2 = __half2half2(__float2half(scale));++ int byte_base = sf << 3;++ uint32_t a4_0 = __ldg(reinterpret_cast<const uint32_t*>(&row_a[byte_base]));+ uint32_t b4_0 = __ldg(reinterpret_cast<const uint32_t*>(&batch_b[byte_base]));+ uint32_t a4_1 = __ldg(reinterpret_cast<const uint32_t*>(&row_a[byte_base + 4]));+ uint32_t b4_1 = __ldg(reinterpret_cast<const uint32_t*>(&batch_b[byte_base + 4]));++ __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);+ __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);++ float2 f0 = __half22float2(acc_h2_0);+ float2 f1 = __half22float2(acc_h2_1);+ acc += f0.x + f0.y + f1.x + f1.y;}-- if (!valid) return;-- // Warp shuffle reduction- #pragma unroll- for (int offset = 16; offset > 0; offset /= 2) {- sum += __shfl_down_sync(0xffffffff, sum, offset);++ return acc;+ }++ __device__ __forceinline__ float compute_remainder_smem(+ const uint8_t* sh_a,+ const uint8_t* sh_b,+ const uint8_t* sh_sfa,+ const uint8_t* sh_sfb,+ int remainder_scales,+ int tid+ ) {+ float acc = 0.0f;++ #pragma unroll 4+ for (int sf = tid; sf < remainder_scales; sf += BLOCK_SIZE) {+ float scale = decode_fp8(static_cast<int8_t>(sh_sfa[sf])) *+ decode_fp8(static_cast<int8_t>(sh_sfb[sf]));+ __half2 scale_h2 = __half2half2(__float2half(scale));++ int byte_base = sf << 3;++ uint32_t a4_0 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base]);+ uint32_t b4_0 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base]);+ uint32_t a4_1 = *reinterpret_cast<const uint32_t*>(&sh_a[byte_base + 4]);+ uint32_t b4_1 = *reinterpret_cast<const uint32_t*>(&sh_b[byte_base + 4]);++ __half2 acc_h2_0 = dot_scaled_4bytes(a4_0, b4_0, scale_h2);+ __half2 acc_h2_1 = dot_scaled_4bytes(a4_1, b4_1, scale_h2);++ float2 f0 = __half22float2(acc_h2_0);+ float2 f1 = __half22float2(acc_h2_1);+ acc += f0.x + f0.y + f1.x + f1.y;}-- if (tidx == 0) {- C[row] = __float2half(sum);- }++ return acc;}- // Kernel for l > 1 with double-buffered B- __global__ __launch_bounds__(BLOCK_SIZE, 2)- void nvfp4_gemv_kernel_batched(- half* __restrict__ C,- const uint8_t* __restrict__ A,- const uint8_t* __restrict__ B,- const uint8_t* __restrict__ SFA,- const uint8_t* __restrict__ SFB,- int m, int k_packed, int k_sf, int l, int n_pad+ // ============================================================================+ // Main GEMV kernel - one block per (row, batch)+ // ============================================================================+ __global__ __launch_bounds__(BLOCK_SIZE)+ void gemv_nvfp4_kernel(+ const int8_t* __restrict__ a,+ const int8_t* __restrict__ b,+ const int8_t* __restrict__ sfa,+ const int8_t* __restrict__ sfb,+ half* __restrict__ c,+ int M, int K, int L,+ int N_rows) {- __shared__ __align__(16) uint8_t smem_B[2][K_TILE_SIZE];- __shared__ __align__(16) uint8_t smem_SFB[2][SF_TILE_SIZE];-- const int tidx = threadIdx.x;- const int tidy = threadIdx.y;- const int tid = tidy * THREADS_PER_K + tidx;- const int row = blockIdx.x * THREADS_PER_M + tidy;- const int batch = blockIdx.z;- const bool valid = (row < m);-- // Batch offsets- const size_t a_batch = batch * (size_t)(m * k_packed);- const size_t b_batch = batch * (size_t)(n_pad * k_packed);- const size_t sfa_batch = batch * (size_t)(m * k_sf);- const size_t sfb_batch = batch * (size_t)(n_pad * k_sf);-- const int num_tiles = (k_packed + K_TILE_SIZE - 1) / K_TILE_SIZE;-- float sum = 0.0f;-- // Load first tile- if (num_tiles > 0) {- if (tid < 32) {- int k_off = tid * 16;- if (k_off < k_packed) {- uint32_t dst = __cvta_generic_to_shared(&smem_B[0][tid * 16]);- ASYNC_COPY_16(dst, B + b_batch + k_off);- }+ // Double-buffered shared memory+ __shared__ __align__(16) uint8_t sh_a[NUM_BUFFERS][BYTES_PER_TILE];+ __shared__ __align__(16) uint8_t sh_b[NUM_BUFFERS][BYTES_PER_TILE];+ __shared__ __align__(16) uint8_t sh_sfa[NUM_BUFFERS][SCALES_PER_TILE];+ __shared__ __align__(16) uint8_t sh_sfb[NUM_BUFFERS][SCALES_PER_TILE];+ __shared__ float smem_acc[BLOCK_SIZE / 32];++ const int m = blockIdx.x;+ const int l = blockIdx.y;+ const int tid = threadIdx.x;++ if (m >= M) return;++ // Dimension calculations+ const int K_bytes = K / 2;+ const int K_sf = K / 16;++ // Batch strides+ 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;++ // Pointers for this row and batch+ const uint8_t* row_a = reinterpret_cast<const uint8_t*>(a) + l * a_batch_stride + m * K_bytes;+ const uint8_t* batch_b = reinterpret_cast<const uint8_t*>(b) + l * b_batch_stride;+ const uint8_t* row_sfa = reinterpret_cast<const uint8_t*>(sfa) + l * sfa_batch_stride + m * K_sf;+ const uint8_t* batch_sfb = reinterpret_cast<const uint8_t*>(sfb) + l * sfb_batch_stride;++ // Tile counts+ const int tile_count = K_bytes / BYTES_PER_TILE;+ const int remainder_start = tile_count * BYTES_PER_TILE;++ float acc = 0.0f;++ // Async copy macros+ #define ASYNC_COPY_16(dst, src) \+ asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" :: "r"(dst), "l"(src))+ #define ASYNC_COPY_4(dst, src) \+ asm volatile("cp.async.ca.shared.global [%0], [%1], 4;" :: "r"(dst), "l"(src))++ // Lambda for issuing tile async copies+ auto issue_tile_async = [&](int b_idx, int tile) {+ const int base_byte = tile * BYTES_PER_TILE;+ const int base_sf = tile * SCALES_PER_TILE;+ const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][0]);+ const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);+ const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][0]);+ const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);++ // Copy A and B tiles (BYTES_PER_TILE bytes each)+ for (int i = tid * 16; i < BYTES_PER_TILE; i += BLOCK_SIZE * 16) {+ ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);+ ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);}- if (tid < 4) {- int sf_off = tid * 16;- if (sf_off < k_sf) {- uint32_t dst = __cvta_generic_to_shared(&smem_SFB[0][tid * 16]);- ASYNC_COPY_16(dst, SFB + sfb_batch + sf_off);- }+ // Copy scale factors (SCALES_PER_TILE bytes each)+ for (int i = tid * 4; i < SCALES_PER_TILE; i += BLOCK_SIZE * 4) {+ ASYNC_COPY_4(sh_sfa_base + i, row_sfa + base_sf + i);+ ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + base_sf + i);}- ASYNC_COMMIT();- }-- for (int tile = 0; tile < num_tiles; tile++) {- int curr = tile & 1;- int next = (tile + 1) & 1;- int next_tile = tile + 1;-- // Prefetch next tile- if (next_tile < num_tiles) {- int k_base_next = next_tile * K_TILE_SIZE;- if (tid < 32) {- int k_off = k_base_next + tid * 16;- if (k_off < k_packed) {- uint32_t dst = __cvta_generic_to_shared(&smem_B[next][tid * 16]);- ASYNC_COPY_16(dst, B + b_batch + k_off);- }- }- if (tid < 4) {- int sf_off = k_base_next / 8 + tid * 16;- if (sf_off < k_sf) {- uint32_t dst = __cvta_generic_to_shared(&smem_SFB[next][tid * 16]);- ASYNC_COPY_16(dst, SFB + sfb_batch + sf_off);- }- }- ASYNC_COMMIT();+ asm volatile("cp.async.commit_group;");+ };++ auto issue_remainder_async = [&](int b_idx, int rem_sf_start, int total_K_sf) {+ const int base_byte = rem_sf_start << 3; // rem_sf_start * 8+ const uint32_t sh_a_base = __cvta_generic_to_shared(&sh_a[b_idx][0]);+ const uint32_t sh_b_base = __cvta_generic_to_shared(&sh_b[b_idx][0]);+ const uint32_t sh_sfa_base = __cvta_generic_to_shared(&sh_sfa[b_idx][0]);+ const uint32_t sh_sfb_base = __cvta_generic_to_shared(&sh_sfb[b_idx][0]);++ int remainder_bytes = (total_K_sf - rem_sf_start) << 3; // * 8+ int remainder_scales = total_K_sf - rem_sf_start;++ for (int i = tid * 16; i < remainder_bytes; i += BLOCK_SIZE * 16) {+ ASYNC_COPY_16(sh_a_base + i, row_a + base_byte + i);+ ASYNC_COPY_16(sh_b_base + i, batch_b + base_byte + i);}-- ASYNC_WAIT_ALL();+ for (int i = tid * 4; i < remainder_scales; i += BLOCK_SIZE * 4) {+ ASYNC_COPY_4(sh_sfa_base + i, row_sfa + rem_sf_start + i);+ ASYNC_COPY_4(sh_sfb_base + i, batch_sfb + rem_sf_start + i);+ }+ asm volatile("cp.async.commit_group;");+ };++ // Main loop+ int remainder_sf_start = remainder_start / 8; // Convert bytes to scale factor index+ bool has_remainder = (remainder_sf_start < K_sf);+ int buf = 0;++ if (tile_count > 0) {+ // Load first tile+ issue_tile_async(0, 0);+ asm volatile("cp.async.wait_group 0;");__syncthreads();-- if (valid) {- int k_base = tile * K_TILE_SIZE;- int k_local = tidx * 16;- int k_global = k_base + k_local;-- if (k_global + 16 <= k_packed) {- uint4 a_vec = *reinterpret_cast<const uint4*>(A + a_batch + row * k_packed + k_global);- uint4 b_vec = *reinterpret_cast<const uint4*>(&smem_B[curr][k_local]);-- const uint8_t* a_bytes = reinterpret_cast<const uint8_t*>(&a_vec);- const uint8_t* b_bytes = reinterpret_cast<const uint8_t*>(&b_vec);-- int sf_local = k_local / 8;- int sf_global = k_global / 8;-- float sfa0 = decode_fp8(SFA[sfa_batch + row * k_sf + sf_global]);- float sfa1 = decode_fp8(SFA[sfa_batch + row * k_sf + sf_global + 1]);- float sfb0 = decode_fp8(smem_SFB[curr][sf_local]);- float sfb1 = decode_fp8(smem_SFB[curr][sf_local + 1]);- float scale0 = sfa0 * sfb0;- float scale1 = sfa1 * sfb1;-- __half2 acc0 = __float2half2_rn(0.0f);- __half2 acc1 = __float2half2_rn(0.0f);-- #pragma unroll- for (int i = 0; i < 8; i++) {- acc0 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc0);- }- sum += (__half2float(__low2half(acc0)) + __half2float(__high2half(acc0))) * scale0;-- #pragma unroll- for (int i = 8; i < 16; i++) {- acc1 = __hfma2(decode_fp4x2(a_bytes[i]), decode_fp4x2(b_bytes[i]), acc1);- }- sum += (__half2float(__low2half(acc1)) + __half2float(__high2half(acc1))) * scale1;++ for (int tile = 0; tile < tile_count; ++tile) {+ if (tile + 1 < tile_count) {+ // Prefetch next tile+ issue_tile_async(buf ^ 1, tile + 1);+ } else if (has_remainder) {+ // On last tile: prefetch remainder+ issue_remainder_async(buf ^ 1, remainder_sf_start, K_sf);}++ // Compute current tile+ acc += compute_tile(+ sh_a[buf],+ sh_b[buf],+ sh_sfa[buf],+ sh_sfb[buf],+ tid+ );++ if (tile + 1 < tile_count || has_remainder) {+ asm volatile("cp.async.wait_group 0;");+ __syncthreads();+ buf ^= 1;+ }}-- __syncthreads();}-- if (!valid) return;-- #pragma unroll- for (int offset = 16; offset > 0; offset /= 2) {- sum += __shfl_down_sync(0xffffffff, sum, offset);++ if (has_remainder) {+ if (tile_count > 0) {+ int remainder_scales = K_sf - remainder_sf_start;+ acc += compute_remainder_smem(+ sh_a[buf],+ sh_b[buf],+ sh_sfa[buf],+ sh_sfb[buf],+ remainder_scales,+ tid+ );+ } else {+ // No tiles - load remainder directly from global memory+ acc += compute_remainder(+ row_a,+ batch_b,+ row_sfa,+ batch_sfb,+ remainder_sf_start,+ K_sf,+ tid+ );+ }}++ #undef ASYNC_COPY_16+ #undef ASYNC_COPY_4++ // ========================================================================+ // Warp-level reduction+ // ========================================================================+ float warp_sum = acc;+ warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 16);+ warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 8);+ warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 4);+ warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 2);+ warp_sum += __shfl_down_sync(0xffffffff, warp_sum, 1);++ // ========================================================================+ // Block-level reduction (only 1 warp, so just write)+ // ========================================================================+ const int warp_id = tid >> 5;+ const int lane = tid & 31;- if (tidx == 0) {- C[batch * m + row] = __float2half(sum);+ if (lane == 0) {+ smem_acc[warp_id] = warp_sum;}+ __syncthreads();++ if (warp_id == 0) {+ float block_sum = (lane < (BLOCK_SIZE >> 5)) ? smem_acc[lane] : 0.0f;+ block_sum += __shfl_down_sync(0xffffffff, block_sum, 16);+ block_sum += __shfl_down_sync(0xffffffff, block_sum, 8);+ block_sum += __shfl_down_sync(0xffffffff, block_sum, 4);+ block_sum += __shfl_down_sync(0xffffffff, block_sum, 2);+ block_sum += __shfl_down_sync(0xffffffff, block_sum, 1);++ // ====================================================================+ // Final output write+ // ====================================================================+ if (lane == 0) {+ size_t c_idx = (size_t)m + (size_t)l * M;+ c[c_idx] = __float2half(block_sum);+ }+ }}+ // ============================================================================+ // Host wrapper+ // ============================================================================void run_nvfp4_gemv(torch::Tensor C,torch::Tensor A,⋯ 2 unchanged linestorch::Tensor SFB,int m, int k, int l, int n_pad) {- int k_packed = k / 2;- int k_sf = k / SF_VEC_SIZE;-- half* c_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());- const uint8_t* a_ptr = A.data_ptr<uint8_t>();- const uint8_t* b_ptr = B.data_ptr<uint8_t>();- const uint8_t* sfa_ptr = SFA.data_ptr<uint8_t>();- const uint8_t* sfb_ptr = SFB.data_ptr<uint8_t>();-- int blocks_m = (m + THREADS_PER_M - 1) / THREADS_PER_M;- dim3 block(THREADS_PER_K, THREADS_PER_M, 1);+ dim3 grid(m, l);+ dim3 block(BLOCK_SIZE);- if (l == 1) {- dim3 grid(blocks_m, 1, 1);- nvfp4_gemv_kernel_single<<<grid, block>>>(- c_ptr, a_ptr, b_ptr, sfa_ptr, sfb_ptr,- m, k_packed, k_sf, n_pad- );- } else {- dim3 grid(blocks_m, 1, l);- nvfp4_gemv_kernel_batched<<<grid, block>>>(- c_ptr, a_ptr, b_ptr, sfa_ptr, sfb_ptr,- m, k_packed, k_sf, l, n_pad- );- }+ gemv_nvfp4_kernel<<<grid, block>>>(+ reinterpret_cast<const int8_t*>(A.data_ptr()),+ reinterpret_cast<const int8_t*>(B.data_ptr()),+ reinterpret_cast<const int8_t*>(SFA.data_ptr()),+ reinterpret_cast<const int8_t*>(SFB.data_ptr()),+ reinterpret_cast<half*>(C.data_ptr<at::Half>()),+ m, k, l, n_pad+ );}-- #undef ASYNC_COPY_16- #undef ASYNC_COMMIT- #undef ASYNC_WAIT_ALL'''CPP_SRC = r'''⋯ 15 unchanged linesglobal _cuda_moduleif _cuda_module is None:_cuda_module = load_inline(- name='nvfp4_gemv_enhanced_v1',+ name='nvfp4_gemv_async_v2',cpp_sources=CPP_SRC,cuda_sources=CUDA_SRC,functions=['run_nvfp4_gemv'],extra_cuda_cflags=[- '-O3',+ '-O3','--use_fast_math','-std=c++17','-gencode=arch=compute_100a,code=sm_100a'
scrolls · 694 diff lines total
Best evidence level for this revision: reported
JSON