submission 74473
s.am._ · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 436 lines, June 9 Researcher Reciprocity License v1.0.
less_naive_v2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-74473?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:2a226c1dab9e5a3c914a2d450ae274858e8768a9d03d94635f4b8bcdaf8dbc2f
license declaredunknown
license concludedunknown
authorss.am._
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;Kernel source
less_naive_v2.py436 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
# ---- C++ stub: declare the function so load_inline can bind it ----
gemv_cpp = r"""
#include <torch/extension.h>
// Forward declaration so PyTorch can bind it (definition is in the CUDA source).
torch::Tensor nvfp4_gemv_v2(torch::Tensor A,
torch::Tensor B,
torch::Tensor C,
torch::Tensor SFA,
torch::Tensor SFB);
"""
# ---- CUDA source: struct, kernel, launcher, and Python-facing wrapper ----
gemv_cuda = r"""
#include <assert.h>
#include <cuda.h>
#include <stdio.h>
#include <cuda_runtime.h>
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cuda_fp4.h>
#include <cuda_bf16.h>
#include <cuda_fp8.h>
// ---- gemv.h ----
struct Gemv_params {
using index_t = uint64_t;
int b, m, k, real_k;
void *__restrict__ a_ptr;
void *__restrict__ b_ptr;
void *__restrict__ sfa_ptr;
void *__restrict__ sfb_ptr;
void *__restrict__ o_ptr;
index_t a_batch_stride;
index_t b_batch_stride;
index_t sfa_batch_stride;
index_t sfb_batch_stride;
index_t o_batch_stride;
index_t a_row_stride;
index_t b_row_stride;
index_t sfa_row_stride;
index_t sfb_row_stride;
index_t o_row_stride;
};
static constexpr int ROWS_PER_BLOCK = 8;
static constexpr int THREADS_PER_ROW = 16;
static constexpr int BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW; // 128
__device__ __forceinline__ void load_block_16x2fp4(
const __nv_fp4x2_e2m1* rowA,
const __nv_fp4x2_e2m1* vecB,
const uint16_t* rowS_u16,
const uint16_t* vecS_u16,
int elem_base,
int block_base,
uint64_t (&a_regs)[2],
uint64_t (&b_regs)[2],
uint16_t &sfa_regs,
uint16_t &sfb_regs)
{
uint64_t rowA_addr = reinterpret_cast<uint64_t>(rowA + elem_base);
uint64_t vecB_addr = reinterpret_cast<uint64_t>(vecB + elem_base);
uint64_t rowS_addr = reinterpret_cast<uint64_t>(rowS_u16 + block_base);
uint64_t vecS_addr = reinterpret_cast<uint64_t>(vecS_u16 + block_base);
asm volatile(
"ld.global.u64.v2 {%0, %1}, [%4];\n\t"
"ld.global.u64.v2 {%2, %3}, [%5];\n\t"
: "=l"(a_regs[0]), "=l"(a_regs[1]),
"=l"(b_regs[0]), "=l"(b_regs[1])
: "l"(rowA_addr), "l"(vecB_addr)
);
asm volatile(
"ld.global.u16 %0, [%2];\n\t"
"ld.global.u16 %1, [%3];\n\t"
: "=h"(sfa_regs), "=h"(sfb_regs)
: "l"(rowS_addr), "l"(vecS_addr)
);
}
__device__ __forceinline__ __half block_scaled_fma_16x2fp4(
const uint64_t (&a_regs)[2],
const uint64_t (&b_regs)[2],
uint16_t sfa_regs,
uint16_t sfb_regs)
{
uint32_t const* a_regs_packed = reinterpret_cast<uint32_t const*>(&a_regs);
uint32_t const* b_regs_packed = reinterpret_cast<uint32_t const*>(&b_regs);
uint16_t out_half_bits;
asm volatile(
"{\n"
".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\n"
".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\n"
".reg .b8 byte0_8, byte0_9, byte0_10, byte0_11;\n"
".reg .b8 byte0_12, byte0_13, byte0_14, byte0_15;\n"
".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\n"
".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\n"
".reg .b8 byte1_8, byte1_9, byte1_10, byte1_11;\n"
".reg .b8 byte1_12, byte1_13, byte1_14, byte1_15;\n"
".reg .f16x2 accum_0, accum_1, accum_2, accum_3;\n"
".reg .f16x2 accum_4, accum_5, accum_6, accum_7;\n"
".reg .f16x2 accum_8, accum_9, accum_10, accum_11;\n"
".reg .f16x2 accum_12, accum_13, accum_14, accum_15;\n"
".reg .f16x2 sfa_f16x2;\n"
".reg .f16x2 sfb_f16x2;\n"
".reg .f16x2 sf_f16x2;\n"
".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\n"
".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\n"
".reg .f16x2 cvt_0_8, cvt_0_9, cvt_0_10, cvt_0_11;\n"
".reg .f16x2 cvt_0_12, cvt_0_13, cvt_0_14, cvt_0_15;\n"
".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\n"
".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\n"
".reg .f16x2 cvt_1_8, cvt_1_9, cvt_1_10, cvt_1_11;\n"
".reg .f16x2 cvt_1_12, cvt_1_13, cvt_1_14, cvt_1_15;\n"
".reg .f16 result_f16, lane0, lane1;\n"
".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\n"
"cvt.rn.f16x2.e4m3x2 sfa_f16x2, %9;\n"
"cvt.rn.f16x2.e4m3x2 sfb_f16x2, %10;\n"
"mov.b32 accum_0, 0;\n"
"mov.b32 accum_1, 0;\n"
"mov.b32 accum_2, 0;\n"
"mov.b32 accum_3, 0;\n"
"mov.b32 accum_4, 0;\n"
"mov.b32 accum_5, 0;\n"
"mov.b32 accum_6, 0;\n"
"mov.b32 accum_7, 0;\n"
"mov.b32 accum_8, 0;\n"
"mov.b32 accum_9, 0;\n"
"mov.b32 accum_10, 0;\n"
"mov.b32 accum_11, 0;\n"
"mov.b32 accum_12, 0;\n"
"mov.b32 accum_13, 0;\n"
"mov.b32 accum_14, 0;\n"
"mov.b32 accum_15, 0;\n"
"mul.rn.f16x2 sf_f16x2, sfa_f16x2, sfb_f16x2;\n"
"mov.b32 {lane0, lane1}, sf_f16x2;\n"
"mov.b32 mul_f16x2_0, {lane0, lane0};\n"
"mov.b32 mul_f16x2_1, {lane1, lane1};\n"
"mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %1;\n"
"mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %2;\n"
"mov.b32 {byte0_8, byte0_9, byte0_10, byte0_11}, %3;\n"
"mov.b32 {byte0_12, byte0_13, byte0_14, byte0_15}, %4;\n"
"mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %5;\n"
"mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %6;\n"
"mov.b32 {byte1_8, byte1_9, byte1_10, byte1_11}, %7;\n"
"mov.b32 {byte1_12, byte1_13, byte1_14, byte1_15}, %8;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_0, byte0_0;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_1, byte0_1;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_2, byte0_2;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_3, byte0_3;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_4, byte0_4;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_5, byte0_5;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_6, byte0_6;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_7, byte0_7;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_8, byte0_8;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_9, byte0_9;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_10, byte0_10;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_11, byte0_11;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_12, byte0_12;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_13, byte0_13;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_14, byte0_14;\n"
"cvt.rn.f16x2.e2m1x2 cvt_0_15, byte0_15;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_0, byte1_0;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_1, byte1_1;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_2, byte1_2;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_3, byte1_3;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_4, byte1_4;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_5, byte1_5;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_6, byte1_6;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_7, byte1_7;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_8, byte1_8;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_9, byte1_9;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_10, byte1_10;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_11, byte1_11;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_12, byte1_12;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_13, byte1_13;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_14, byte1_14;\n"
"cvt.rn.f16x2.e2m1x2 cvt_1_15, byte1_15;\n"
"fma.rn.f16x2 accum_0, cvt_0_0, cvt_1_0, accum_0;\n"
"fma.rn.f16x2 accum_1, cvt_0_1, cvt_1_1, accum_1;\n"
"fma.rn.f16x2 accum_2, cvt_0_2, cvt_1_2, accum_2;\n"
"fma.rn.f16x2 accum_3, cvt_0_3, cvt_1_3, accum_3;\n"
"fma.rn.f16x2 accum_4, cvt_0_4, cvt_1_4, accum_4;\n"
"fma.rn.f16x2 accum_5, cvt_0_5, cvt_1_5, accum_5;\n"
"fma.rn.f16x2 accum_6, cvt_0_6, cvt_1_6, accum_6;\n"
"fma.rn.f16x2 accum_7, cvt_0_7, cvt_1_7, accum_7;\n"
"fma.rn.f16x2 accum_8, cvt_0_8, cvt_1_8, accum_8;\n"
"fma.rn.f16x2 accum_9, cvt_0_9, cvt_1_9, accum_9;\n"
"fma.rn.f16x2 accum_10, cvt_0_10, cvt_1_10, accum_10;\n"
"fma.rn.f16x2 accum_11, cvt_0_11, cvt_1_11, accum_11;\n"
"fma.rn.f16x2 accum_12, cvt_0_12, cvt_1_12, accum_12;\n"
"fma.rn.f16x2 accum_13, cvt_0_13, cvt_1_13, accum_13;\n"
"fma.rn.f16x2 accum_14, cvt_0_14, cvt_1_14, accum_14;\n"
"fma.rn.f16x2 accum_15, cvt_0_15, cvt_1_15, accum_15;\n"
"add.rn.f16x2 accum_0, accum_0, accum_1;\n"
"add.rn.f16x2 accum_2, accum_2, accum_3;\n"
"add.rn.f16x2 accum_4, accum_4, accum_5;\n"
"add.rn.f16x2 accum_6, accum_6, accum_7;\n"
"add.rn.f16x2 accum_8, accum_8, accum_9;\n"
"add.rn.f16x2 accum_10, accum_10, accum_11;\n"
"add.rn.f16x2 accum_12, accum_12, accum_13;\n"
"add.rn.f16x2 accum_14, accum_14, accum_15;\n"
"add.rn.f16x2 accum_0, accum_0, accum_2;\n"
"add.rn.f16x2 accum_4, accum_4, accum_6;\n"
"add.rn.f16x2 accum_8, accum_8, accum_10;\n"
"add.rn.f16x2 accum_12, accum_12, accum_14;\n"
"add.rn.f16x2 accum_0, accum_0, accum_4;\n"
"add.rn.f16x2 accum_8, accum_8, accum_12;\n"
"mul.rn.f16x2 accum_0, mul_f16x2_0, accum_0;\n"
"mul.rn.f16x2 accum_8, mul_f16x2_1, accum_8;\n"
"add.rn.f16x2 accum_0, accum_0, accum_8;\n"
"mov.b32 {lane0, lane1}, accum_0;\n"
"add.rn.f16 result_f16, lane0, lane1;\n"
"mov.b16 %0, result_f16;\n"
"}\n"
: "=h"(out_half_bits)
: "r"(a_regs_packed[0]), "r"(a_regs_packed[1]),
"r"(a_regs_packed[2]), "r"(a_regs_packed[3]),
"r"(b_regs_packed[0]), "r"(b_regs_packed[1]),
"r"(b_regs_packed[2]), "r"(b_regs_packed[3]),
"h"(sfa_regs), "h"(sfb_regs)
: "memory"
);
union { uint16_t u; __half h; } conv;
conv.u = out_half_bits;
return conv.h;
}
__global__ void __launch_bounds__(BLOCK_SIZE, 8)
gemv_kernel(const __grid_constant__ Gemv_params params)
{
const int tid = threadIdx.x;
const int rib = tid / THREADS_PER_ROW;
const int lane = tid % THREADS_PER_ROW;
const int batch = blockIdx.z;
const int row = blockIdx.x * ROWS_PER_BLOCK + rib;
const size_t A_batch_base = static_cast<size_t>(batch) * params.a_batch_stride;
const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;
const size_t B_batch_base = static_cast<size_t>(batch) * params.b_batch_stride;
const size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;
const size_t C_batch_base = static_cast<size_t>(batch) * params.o_batch_stride;
const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base + row * params.a_row_stride;
const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;
const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;
const __nv_fp8_e4m3* vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;
const uint16_t* rowS_u16 = reinterpret_cast<const uint16_t*>(rowS);
const uint16_t* vecS_u16 = reinterpret_cast<const uint16_t*>(vecS);
float sum = 0.f;
const int full_block = THREADS_PER_ROW * 16;
if (params.k >= full_block) {
int iters = params.k / full_block;
for (int idx = 0; idx < iters; ++idx) {
int block_base = idx * THREADS_PER_ROW + lane;
int elem_base = block_base * 16;
uint64_t a_regs[2], b_regs[2];
uint16_t sfa_regs, sfb_regs;
load_block_16x2fp4(
rowA, vecB,
rowS_u16, vecS_u16,
elem_base, block_base,
a_regs, b_regs,
sfa_regs, sfb_regs);
__half h = block_scaled_fma_16x2fp4(a_regs, b_regs, sfa_regs, sfb_regs);
sum += __half2float(h);
}
} else {
int iters = params.k / THREADS_PER_ROW / 8;
for (int idx = 0; idx < iters; ++idx) {
int base = idx * 16;
int base_id = (idx * THREADS_PER_ROW + lane) * 8;
__nv_fp8_storage_t sfa_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&rowS[base + lane]);
__nv_fp8_storage_t sfb_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&vecS[base + lane]);
__half sfa = __nv_cvt_fp8_to_halfraw(sfa_storage, __NV_E4M3);
__half sfb = __nv_cvt_fp8_to_halfraw(sfb_storage, __NV_E4M3);
__half scale = __hmul(sfa, sfb);
__half2 acc = __float2half2_rn(0.0f);
#pragma unroll
for (int i = 0; i < 8; ++i) {
const int id = base_id + i;
__nv_fp4x2_storage_t a_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&rowA[id]);
__nv_fp4x2_storage_t b_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&vecB[id]);
__half2_raw a_raw = __nv_cvt_fp4x2_to_halfraw2(a_storage, __NV_E2M1);
__half2_raw b_raw = __nv_cvt_fp4x2_to_halfraw2(b_storage, __NV_E2M1);
const __half2 a_h2 = __half2(a_raw);
const __half2 b_h2 = __half2(b_raw);
acc = __hfma2(a_h2, b_h2, acc);
}
__half fin = __hadd(__low2half(acc), __high2half(acc));
__half h = __hmul(fin, scale);
sum += __half2float(h);
}
}
unsigned mask = 0xffffffffu;
sum += __shfl_down_sync(mask, sum, 8, 16);
sum += __shfl_down_sync(mask, sum, 4, 16);
sum += __shfl_down_sync(mask, sum, 2, 16);
sum += __shfl_down_sync(mask, sum, 1, 16);
if (lane == 0) {
__half* out = (__half*)params.o_ptr + C_batch_base + row;
out[0] = __float2half(sum);
}
}
static inline void launch_kernel(Gemv_params ¶ms, cudaStream_t stream)
{
const int grid_x = (params.m + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK;
dim3 grid(grid_x, 1, params.b);
dim3 block(BLOCK_SIZE, 1, 1);
gemv_kernel<<<grid, block, 0, stream>>>(params);
}
// Python-facing function: sets up params and launches
torch::Tensor nvfp4_gemv_v2(torch::Tensor A,
torch::Tensor B,
torch::Tensor C,
torch::Tensor SFA,
torch::Tensor SFB)
{
auto sizes = A.sizes();
const int64_t M = sizes[0];
const int64_t K = sizes[1];
const int64_t L = sizes[2];
Gemv_params params{};
params.b = static_cast<int>(L);
params.m = static_cast<int>(M);
params.k = static_cast<int>(K);
params.real_k = static_cast<int>(K * 2);
params.a_ptr = A.data_ptr();
params.b_ptr = B.data_ptr();
params.sfa_ptr= SFA.data_ptr();
params.sfb_ptr= SFB.data_ptr();
params.o_ptr = C.data_ptr();
params.a_batch_stride = static_cast<uint64_t>(A.stride(2));
params.b_batch_stride = static_cast<uint64_t>(B.stride(2));
params.sfa_batch_stride= static_cast<uint64_t>(SFA.stride(2));
params.sfb_batch_stride= static_cast<uint64_t>(SFB.stride(2));
params.o_batch_stride = static_cast<uint64_t>(C.stride(2));
params.a_row_stride = static_cast<uint64_t>(A.stride(0));
params.b_row_stride = static_cast<uint64_t>(B.stride(0));
params.sfa_row_stride= static_cast<uint64_t>(SFA.stride(0));
params.sfb_row_stride= static_cast<uint64_t>(SFB.stride(0));
params.o_row_stride = static_cast<uint64_t>(C.stride(0));
auto stream = at::cuda::getCurrentCUDAStream().stream();
launch_kernel(params, stream);
cudaError_t err = cudaGetLastError();
return C;
}
"""
# ---- build the module ----
nvfp4_module = load_inline(
name="nvfp4_gemv",
cpp_sources=[gemv_cpp],
cuda_sources=[gemv_cuda],
functions=["nvfp4_gemv_v2"], # this exposes the function to Python
extra_cuda_cflags=[
"-std=c++17",
"-gencode=arch=compute_100a,code=sm_100a",
"--ptxas-options=--gpu-name=sm_100a",
"-O3",
"-w",
"--use_fast_math",
"-allow-unsupported-compiler",
],
extra_ldflags=["-lcuda", "-lcublas"],
verbose=True,
)
def custom_kernel(data: input_t) -> output_t:
a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data
return nvfp4_module.nvfp4_gemv_v2(a, b, c, sfa, sfb)
scrolls · 436 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 71977.
⋯ 6 unchanged lines#include <torch/extension.h>// Forward declaration so PyTorch can bind it (definition is in the CUDA source).- torch::Tensor nvfp4_gemv_v1(torch::Tensor A,+ torch::Tensor nvfp4_gemv_v2(torch::Tensor A,torch::Tensor B,torch::Tensor C,torch::Tensor SFA,⋯ 40 unchanged linesindex_t o_row_stride;};- // ---- gemv_v1.cu ----static constexpr int ROWS_PER_BLOCK = 8;static constexpr int THREADS_PER_ROW = 16;static constexpr int BLOCK_SIZE = ROWS_PER_BLOCK * THREADS_PER_ROW; // 128+ __device__ __forceinline__ void load_block_16x2fp4(+ const __nv_fp4x2_e2m1* rowA,+ const __nv_fp4x2_e2m1* vecB,+ const uint16_t* rowS_u16,+ const uint16_t* vecS_u16,+ int elem_base,+ int block_base,+ uint64_t (&a_regs)[2],+ uint64_t (&b_regs)[2],+ uint16_t &sfa_regs,+ uint16_t &sfb_regs)+ {+ uint64_t rowA_addr = reinterpret_cast<uint64_t>(rowA + elem_base);+ uint64_t vecB_addr = reinterpret_cast<uint64_t>(vecB + elem_base);+ uint64_t rowS_addr = reinterpret_cast<uint64_t>(rowS_u16 + block_base);+ uint64_t vecS_addr = reinterpret_cast<uint64_t>(vecS_u16 + block_base);++ asm volatile(+ "ld.global.u64.v2 {%0, %1}, [%4];\n\t"+ "ld.global.u64.v2 {%2, %3}, [%5];\n\t"+ : "=l"(a_regs[0]), "=l"(a_regs[1]),+ "=l"(b_regs[0]), "=l"(b_regs[1])+ : "l"(rowA_addr), "l"(vecB_addr)+ );++ asm volatile(+ "ld.global.u16 %0, [%2];\n\t"+ "ld.global.u16 %1, [%3];\n\t"+ : "=h"(sfa_regs), "=h"(sfb_regs)+ : "l"(rowS_addr), "l"(vecS_addr)+ );+ }++ __device__ __forceinline__ __half block_scaled_fma_16x2fp4(+ const uint64_t (&a_regs)[2],+ const uint64_t (&b_regs)[2],+ uint16_t sfa_regs,+ uint16_t sfb_regs)+ {+ uint32_t const* a_regs_packed = reinterpret_cast<uint32_t const*>(&a_regs);+ uint32_t const* b_regs_packed = reinterpret_cast<uint32_t const*>(&b_regs);++ uint16_t out_half_bits;++ asm volatile(+ "{\n"+ ".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\n"+ ".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\n"+ ".reg .b8 byte0_8, byte0_9, byte0_10, byte0_11;\n"+ ".reg .b8 byte0_12, byte0_13, byte0_14, byte0_15;\n"+ ".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\n"+ ".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\n"+ ".reg .b8 byte1_8, byte1_9, byte1_10, byte1_11;\n"+ ".reg .b8 byte1_12, byte1_13, byte1_14, byte1_15;\n"++ ".reg .f16x2 accum_0, accum_1, accum_2, accum_3;\n"+ ".reg .f16x2 accum_4, accum_5, accum_6, accum_7;\n"+ ".reg .f16x2 accum_8, accum_9, accum_10, accum_11;\n"+ ".reg .f16x2 accum_12, accum_13, accum_14, accum_15;\n"++ ".reg .f16x2 sfa_f16x2;\n"+ ".reg .f16x2 sfb_f16x2;\n"+ ".reg .f16x2 sf_f16x2;\n"++ ".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\n"+ ".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\n"+ ".reg .f16x2 cvt_0_8, cvt_0_9, cvt_0_10, cvt_0_11;\n"+ ".reg .f16x2 cvt_0_12, cvt_0_13, cvt_0_14, cvt_0_15;\n"+ ".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\n"+ ".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\n"+ ".reg .f16x2 cvt_1_8, cvt_1_9, cvt_1_10, cvt_1_11;\n"+ ".reg .f16x2 cvt_1_12, cvt_1_13, cvt_1_14, cvt_1_15;\n"+ ".reg .f16 result_f16, lane0, lane1;\n"+ ".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\n"++ "cvt.rn.f16x2.e4m3x2 sfa_f16x2, %9;\n"+ "cvt.rn.f16x2.e4m3x2 sfb_f16x2, %10;\n"++ "mov.b32 accum_0, 0;\n"+ "mov.b32 accum_1, 0;\n"+ "mov.b32 accum_2, 0;\n"+ "mov.b32 accum_3, 0;\n"+ "mov.b32 accum_4, 0;\n"+ "mov.b32 accum_5, 0;\n"+ "mov.b32 accum_6, 0;\n"+ "mov.b32 accum_7, 0;\n"+ "mov.b32 accum_8, 0;\n"+ "mov.b32 accum_9, 0;\n"+ "mov.b32 accum_10, 0;\n"+ "mov.b32 accum_11, 0;\n"+ "mov.b32 accum_12, 0;\n"+ "mov.b32 accum_13, 0;\n"+ "mov.b32 accum_14, 0;\n"+ "mov.b32 accum_15, 0;\n"++ "mul.rn.f16x2 sf_f16x2, sfa_f16x2, sfb_f16x2;\n"+ "mov.b32 {lane0, lane1}, sf_f16x2;\n"+ "mov.b32 mul_f16x2_0, {lane0, lane0};\n"+ "mov.b32 mul_f16x2_1, {lane1, lane1};\n"++ "mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %1;\n"+ "mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %2;\n"+ "mov.b32 {byte0_8, byte0_9, byte0_10, byte0_11}, %3;\n"+ "mov.b32 {byte0_12, byte0_13, byte0_14, byte0_15}, %4;\n"+ "mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %5;\n"+ "mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %6;\n"+ "mov.b32 {byte1_8, byte1_9, byte1_10, byte1_11}, %7;\n"+ "mov.b32 {byte1_12, byte1_13, byte1_14, byte1_15}, %8;\n"++ "cvt.rn.f16x2.e2m1x2 cvt_0_0, byte0_0;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_1, byte0_1;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_2, byte0_2;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_3, byte0_3;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_4, byte0_4;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_5, byte0_5;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_6, byte0_6;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_7, byte0_7;\n"++ "cvt.rn.f16x2.e2m1x2 cvt_0_8, byte0_8;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_9, byte0_9;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_10, byte0_10;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_11, byte0_11;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_12, byte0_12;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_13, byte0_13;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_14, byte0_14;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_0_15, byte0_15;\n"++ "cvt.rn.f16x2.e2m1x2 cvt_1_0, byte1_0;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_1, byte1_1;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_2, byte1_2;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_3, byte1_3;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_4, byte1_4;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_5, byte1_5;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_6, byte1_6;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_7, byte1_7;\n"++ "cvt.rn.f16x2.e2m1x2 cvt_1_8, byte1_8;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_9, byte1_9;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_10, byte1_10;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_11, byte1_11;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_12, byte1_12;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_13, byte1_13;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_14, byte1_14;\n"+ "cvt.rn.f16x2.e2m1x2 cvt_1_15, byte1_15;\n"++ "fma.rn.f16x2 accum_0, cvt_0_0, cvt_1_0, accum_0;\n"+ "fma.rn.f16x2 accum_1, cvt_0_1, cvt_1_1, accum_1;\n"+ "fma.rn.f16x2 accum_2, cvt_0_2, cvt_1_2, accum_2;\n"+ "fma.rn.f16x2 accum_3, cvt_0_3, cvt_1_3, accum_3;\n"+ "fma.rn.f16x2 accum_4, cvt_0_4, cvt_1_4, accum_4;\n"+ "fma.rn.f16x2 accum_5, cvt_0_5, cvt_1_5, accum_5;\n"+ "fma.rn.f16x2 accum_6, cvt_0_6, cvt_1_6, accum_6;\n"+ "fma.rn.f16x2 accum_7, cvt_0_7, cvt_1_7, accum_7;\n"++ "fma.rn.f16x2 accum_8, cvt_0_8, cvt_1_8, accum_8;\n"+ "fma.rn.f16x2 accum_9, cvt_0_9, cvt_1_9, accum_9;\n"+ "fma.rn.f16x2 accum_10, cvt_0_10, cvt_1_10, accum_10;\n"+ "fma.rn.f16x2 accum_11, cvt_0_11, cvt_1_11, accum_11;\n"+ "fma.rn.f16x2 accum_12, cvt_0_12, cvt_1_12, accum_12;\n"+ "fma.rn.f16x2 accum_13, cvt_0_13, cvt_1_13, accum_13;\n"+ "fma.rn.f16x2 accum_14, cvt_0_14, cvt_1_14, accum_14;\n"+ "fma.rn.f16x2 accum_15, cvt_0_15, cvt_1_15, accum_15;\n"++ "add.rn.f16x2 accum_0, accum_0, accum_1;\n"+ "add.rn.f16x2 accum_2, accum_2, accum_3;\n"+ "add.rn.f16x2 accum_4, accum_4, accum_5;\n"+ "add.rn.f16x2 accum_6, accum_6, accum_7;\n"+ "add.rn.f16x2 accum_8, accum_8, accum_9;\n"+ "add.rn.f16x2 accum_10, accum_10, accum_11;\n"+ "add.rn.f16x2 accum_12, accum_12, accum_13;\n"+ "add.rn.f16x2 accum_14, accum_14, accum_15;\n"++ "add.rn.f16x2 accum_0, accum_0, accum_2;\n"+ "add.rn.f16x2 accum_4, accum_4, accum_6;\n"+ "add.rn.f16x2 accum_8, accum_8, accum_10;\n"+ "add.rn.f16x2 accum_12, accum_12, accum_14;\n"++ "add.rn.f16x2 accum_0, accum_0, accum_4;\n"+ "add.rn.f16x2 accum_8, accum_8, accum_12;\n"++ "mul.rn.f16x2 accum_0, mul_f16x2_0, accum_0;\n"+ "mul.rn.f16x2 accum_8, mul_f16x2_1, accum_8;\n"++ "add.rn.f16x2 accum_0, accum_0, accum_8;\n"++ "mov.b32 {lane0, lane1}, accum_0;\n"+ "add.rn.f16 result_f16, lane0, lane1;\n"++ "mov.b16 %0, result_f16;\n"+ "}\n"+ : "=h"(out_half_bits)+ : "r"(a_regs_packed[0]), "r"(a_regs_packed[1]),+ "r"(a_regs_packed[2]), "r"(a_regs_packed[3]),+ "r"(b_regs_packed[0]), "r"(b_regs_packed[1]),+ "r"(b_regs_packed[2]), "r"(b_regs_packed[3]),+ "h"(sfa_regs), "h"(sfb_regs)+ : "memory"+ );++ union { uint16_t u; __half h; } conv;+ conv.u = out_half_bits;+ return conv.h;+ }+__global__ void __launch_bounds__(BLOCK_SIZE, 8)gemv_kernel(const __grid_constant__ Gemv_params params){- const int tid = threadIdx.x;- const int rib = tid / THREADS_PER_ROW;- const int lane = tid % THREADS_PER_ROW;- const int batch = blockIdx.z;- const int row = blockIdx.x * ROWS_PER_BLOCK + rib;+ const int tid = threadIdx.x;+ const int rib = tid / THREADS_PER_ROW;+ const int lane = tid % THREADS_PER_ROW;+ const int batch = blockIdx.z;+ const int row = blockIdx.x * ROWS_PER_BLOCK + rib;const size_t A_batch_base = static_cast<size_t>(batch) * params.a_batch_stride;const size_t SFA_batch_base = static_cast<size_t>(batch) * params.sfa_batch_stride;⋯ 1 unchanged linesconst size_t SFB_batch_base = static_cast<size_t>(batch) * params.sfb_batch_stride;const size_t C_batch_base = static_cast<size_t>(batch) * params.o_batch_stride;- const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base + row * params.a_row_stride;- const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;-+ const __nv_fp4x2_e2m1* rowA = static_cast<const __nv_fp4x2_e2m1*>(params.a_ptr) + A_batch_base + row * params.a_row_stride;+ const __nv_fp8_e4m3* rowS = static_cast<const __nv_fp8_e4m3*>(params.sfa_ptr) + SFA_batch_base + row * params.sfa_row_stride;const __nv_fp4x2_e2m1* vecB = static_cast<const __nv_fp4x2_e2m1*>(params.b_ptr) + B_batch_base;- const __nv_fp8_e4m3* vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;+ const __nv_fp8_e4m3* vecS = static_cast<const __nv_fp8_e4m3*>(params.sfb_ptr) + SFB_batch_base;+ const uint16_t* rowS_u16 = reinterpret_cast<const uint16_t*>(rowS);+ const uint16_t* vecS_u16 = reinterpret_cast<const uint16_t*>(vecS);+float sum = 0.f;- for (int idx = 0; idx < params.k / THREADS_PER_ROW / 8; ++idx) {- int base = idx * 16;- const int base_id = (idx * THREADS_PER_ROW + lane) * 8;+ const int full_block = THREADS_PER_ROW * 16;- __nv_fp8_storage_t sfa_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&rowS[base + lane]);- __nv_fp8_storage_t sfb_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&vecS[base + lane]);- __half sfa = __nv_cvt_fp8_to_halfraw(sfa_storage, __NV_E4M3);- __half sfb = __nv_cvt_fp8_to_halfraw(sfb_storage, __NV_E4M3);- __half scale = __hmul(sfa, sfb);+ if (params.k >= full_block) {+ int iters = params.k / full_block;- __half2 acc = __float2half2_rn(0.0f);+ for (int idx = 0; idx < iters; ++idx) {+ int block_base = idx * THREADS_PER_ROW + lane;+ int elem_base = block_base * 16;- #pragma unroll- for (int i = 0; i < 8; ++i) {- const int id = base_id + i;- __nv_fp4x2_storage_t a_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&rowA[id]);- __nv_fp4x2_storage_t b_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&vecB[id]);- __half2_raw a_raw = __nv_cvt_fp4x2_to_halfraw2(a_storage, __NV_E2M1);- __half2_raw b_raw = __nv_cvt_fp4x2_to_halfraw2(b_storage, __NV_E2M1);+ uint64_t a_regs[2], b_regs[2];+ uint16_t sfa_regs, sfb_regs;- const __half2 a_h2 = __half2(a_raw);- const __half2 b_h2 = __half2(b_raw);+ load_block_16x2fp4(+ rowA, vecB,+ rowS_u16, vecS_u16,+ elem_base, block_base,+ a_regs, b_regs,+ sfa_regs, sfb_regs);- acc = __hfma2(a_h2, b_h2, acc);+ __half h = block_scaled_fma_16x2fp4(a_regs, b_regs, sfa_regs, sfb_regs);+ sum += __half2float(h);}+ } else {+ int iters = params.k / THREADS_PER_ROW / 8;- __half fin = __hadd(__low2half(acc), __high2half(acc));- __half h = __hmul(fin, scale);+ for (int idx = 0; idx < iters; ++idx) {+ int base = idx * 16;+ int base_id = (idx * THREADS_PER_ROW + lane) * 8;- sum += __half2float(h);+ __nv_fp8_storage_t sfa_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&rowS[base + lane]);+ __nv_fp8_storage_t sfb_storage = *reinterpret_cast<const __nv_fp8_storage_t*>(&vecS[base + lane]);+ __half sfa = __nv_cvt_fp8_to_halfraw(sfa_storage, __NV_E4M3);+ __half sfb = __nv_cvt_fp8_to_halfraw(sfb_storage, __NV_E4M3);+ __half scale = __hmul(sfa, sfb);++ __half2 acc = __float2half2_rn(0.0f);++ #pragma unroll+ for (int i = 0; i < 8; ++i) {+ const int id = base_id + i;+ __nv_fp4x2_storage_t a_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&rowA[id]);+ __nv_fp4x2_storage_t b_storage = *reinterpret_cast<const __nv_fp4x2_storage_t*>(&vecB[id]);+ __half2_raw a_raw = __nv_cvt_fp4x2_to_halfraw2(a_storage, __NV_E2M1);+ __half2_raw b_raw = __nv_cvt_fp4x2_to_halfraw2(b_storage, __NV_E2M1);++ const __half2 a_h2 = __half2(a_raw);+ const __half2 b_h2 = __half2(b_raw);++ acc = __hfma2(a_h2, b_h2, acc);+ }++ __half fin = __hadd(__low2half(acc), __high2half(acc));+ __half h = __hmul(fin, scale);+ sum += __half2float(h);+ }}unsigned mask = 0xffffffffu;⋯ 17 unchanged lines}// Python-facing function: sets up params and launches- torch::Tensor nvfp4_gemv_v1(torch::Tensor A,+ torch::Tensor nvfp4_gemv_v2(torch::Tensor A,torch::Tensor B,torch::Tensor C,torch::Tensor SFA,torch::Tensor SFB){- TORCH_CHECK(A.device().is_cuda(), "A must be CUDA");- TORCH_CHECK(B.device().is_cuda(), "B must be CUDA");- TORCH_CHECK(C.device().is_cuda(), "C must be CUDA");- TORCH_CHECK(SFA.device().is_cuda(), "SFA must be CUDA");- TORCH_CHECK(SFB.device().is_cuda(), "SFB must be CUDA");- // Expect A: [M,K,L], B: [K,L], C: [M,1,L] or similar layout using provided strides.auto sizes = A.sizes();- TORCH_CHECK(sizes.size() == 3, "A must be 3D [M,K,L]");const int64_t M = sizes[0];const int64_t K = sizes[1];const int64_t L = sizes[2];⋯ 26 unchanged lineslaunch_kernel(params, stream);cudaError_t err = cudaGetLastError();- TORCH_CHECK(err == cudaSuccess, "CUDA kernel failed: ", cudaGetErrorString(err));return C;}⋯ 4 unchanged linesname="nvfp4_gemv",cpp_sources=[gemv_cpp],cuda_sources=[gemv_cuda],- functions=["nvfp4_gemv_v1"], # this exposes the function to Python+ functions=["nvfp4_gemv_v2"], # this exposes the function to Pythonextra_cuda_cflags=["-std=c++17","-gencode=arch=compute_100a,code=sm_100a",⋯ 10 unchanged linesdef custom_kernel(data: input_t) -> output_t:a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data- return nvfp4_module.nvfp4_gemv_v1(a, b, c, sfa, sfb)+ return nvfp4_module.nvfp4_gemv_v2(a, b, c, sfa, sfb)
scrolls · 382 diff lines total
Best evidence level for this revision: reported
JSON