submission 104932
yue · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 283 lines, June 9 Researcher Reciprocity License v1.0.
submit_v0_ptx3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-104932?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:bef06ee649827f10d1cc1c988437ea8af19edd37d4934d4efd7933de54f42655
license declaredunknown
license concludedunknown
authorsyue
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
- Vectorized decodes using PTX asm for FP4 x8 and FP8 x4shared-memory
__shared__ __half smem_results[ROWS_PER_BLOCK];tile-k = 64
constexpr int TILE_K = 64; // Each thread handles 64 elementsvector-width = float4
float4 A_data[2]; // 2x16 bytes = 32 bytesKernel source
submit_v0_ptx3.py283 lines
import torch
import sys
from torch.utils.cpp_extension import load_inline
from typing import Tuple
from task import input_t, output_t
gemv_cuda_src = """
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cuda_pipeline.h>
#include <cstdint>
// Vectorized decode for FP4 x8 - FIXED VERSION
// cvt.rn.f16x2.e2m1x2 takes .b8 input and produces .b32 output (f16x2)
__device__ __forceinline__ void decode_fp4x8_e2m1_half4(const uint32_t packed, __half2& h0, __half2& h1, __half2& h2, __half2& h3)
{
uint32_t out0, out1, out2, out3;
// PTX allows wider operands for .b8 instruction types
// Input is uint32 containing 4 bytes, output is 4x .b32 registers
asm volatile (
"{"
" .reg .b8 %%b0, %%b1, %%b2, %%b3;\\n"
" mov.b32 {%%b0, %%b1, %%b2, %%b3}, %4;\\n"
" cvt.rn.f16x2.e2m1x2 %0, %%b0;\\n"
" cvt.rn.f16x2.e2m1x2 %1, %%b1;\\n"
" cvt.rn.f16x2.e2m1x2 %2, %%b2;\\n"
" cvt.rn.f16x2.e2m1x2 %3, %%b3;\\n"
"}"
: "=r"(out0), "=r"(out1), "=r"(out2), "=r"(out3)
: "r"(packed)
);
h0 = *reinterpret_cast<const __half2*>(&out0);
h1 = *reinterpret_cast<const __half2*>(&out1);
h2 = *reinterpret_cast<const __half2*>(&out2);
h3 = *reinterpret_cast<const __half2*>(&out3);
}
// Vectorized decode for FP8 x4 - FIXED VERSION
// cvt.rn.f16x2.e4m3x2 takes .b16 input and produces .b32 output (f16x2)
__device__ __forceinline__ void decode_fp8x4_e4m3fn_half4(const uint32_t packed, __half& h0, __half& h1, __half& h2, __half& h3)
{
uint32_t out_low, out_high;
// Input is uint32 containing 2x16-bit values
asm volatile (
"{"
" .reg .b16 %%low, %%high;\\n"
" mov.b32 {%%low, %%high}, %2;\\n"
" cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"
" cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"
"}"
: "=r"(out_low), "=r"(out_high)
: "r"(packed)
);
__half2 h_low = *reinterpret_cast<const __half2*>(&out_low);
__half2 h_high = *reinterpret_cast<const __half2*>(&out_high);
h0 = h_low.x;
h1 = h_low.y;
h2 = h_high.x;
h3 = h_high.y;
}
extern "C" __global__
void block_scaled_gemv_fp4_fp8_fp16_vectorized(
const uint8_t* __restrict__ A, // [l, m, k//2]
const uint8_t* __restrict__ B, // [l, 128, k//2]
const uint8_t* __restrict__ SFA, // [l, m, k//16]
const uint8_t* __restrict__ SFB, // [l, 128, k//16]
__half* __restrict__ C, // [l, m, 1]
int M, int K, int L)
{
constexpr int ROWS_PER_BLOCK = 32;
constexpr int THREADS_PER_ROW = 16; // 16 threads collaborate on each row
constexpr int TILE_K = 64; // Each thread handles 64 elements
constexpr int SF_VEC_SIZE = 16;
const int tidx = threadIdx.x; // 0-15: which K-segment for this row
const int tidy = threadIdx.y; // 0-31: which row
const int block_row = blockIdx.x * ROWS_PER_BLOCK;
const int batch_idx = blockIdx.z;
const int global_row = block_row + tidy;
// Shared memory for coalesced writes
__shared__ __half smem_results[ROWS_PER_BLOCK];
if (global_row >= M) return;
float local_sum = 0.0f;
const int num_k_tiles = K / TILE_K;
const uint8_t* B_base = B + batch_idx * (128 * K / 2);
const uint8_t* SFB_base = SFB + batch_idx * (128 * K / 16);
const uint8_t* A_row = A + batch_idx * (M * K / 2) + global_row * (K / 2);
const uint8_t* SFA_row = SFA + batch_idx * (M * K / 16) + global_row * (K / 16);
// K-PARALLELISM: Each thread handles multiple K-tiles, striding by THREADS_PER_ROW
#pragma unroll 1
for (int tile = tidx; tile < num_k_tiles; tile += THREADS_PER_ROW) {
const int k_offset = tile * TILE_K;
// Vectorized loads
float4 A_data[2]; // 2x16 bytes = 32 bytes
float4 B_data[2];
const float4* A_ptr = reinterpret_cast<const float4*>(A_row + k_offset / 2);
const float4* B_ptr = reinterpret_cast<const float4*>(B_base + k_offset / 2);
A_data[0] = __ldg(A_ptr);
A_data[1] = __ldg(A_ptr + 1);
B_data[0] = __ldg(B_ptr);
B_data[1] = __ldg(B_ptr + 1);
uint8_t* A_tile_data = reinterpret_cast<uint8_t*>(A_data);
uint8_t* B_tile_data = reinterpret_cast<uint8_t*>(B_data);
// Load scale factors (4 bytes per thread)
uint32_t sfa_vec = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + k_offset / 16));
uint32_t sfb_vec = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + k_offset / 16));
// Vectorized decode for all 4 scales at once
__half sfa_scales[4], sfb_scales[4];
decode_fp8x4_e4m3fn_half4(sfa_vec, sfa_scales[0], sfa_scales[1], sfa_scales[2], sfa_scales[3]);
decode_fp8x4_e4m3fn_half4(sfb_vec, sfb_scales[0], sfb_scales[1], sfb_scales[2], sfb_scales[3]);
// COMPUTE with vectorized decode using PTX asm
constexpr int NUM_SF_BLOCKS = TILE_K / SF_VEC_SIZE;
#pragma unroll
for (int sf_block = 0; sf_block < NUM_SF_BLOCKS; sf_block++) {
const int base_k = sf_block * SF_VEC_SIZE;
// Use pre-decoded scales
__half sfa = sfa_scales[sf_block];
__half sfb = sfb_scales[sf_block];
__half scale = __hmul(sfa, sfb);
float block_sum = 0.0f;
// Process 16 elements (8 bytes packed) using two x8 decodes
#pragma unroll
for (int sub = 0; sub < 2; sub++) {
const int byte_idx = base_k / 2 + sub * 4;
const uint32_t a_packed = *reinterpret_cast<const uint32_t*>(A_tile_data + byte_idx);
const uint32_t b_packed = *reinterpret_cast<const uint32_t*>(B_tile_data + byte_idx);
__half2 a0, a1, a2, a3;
__half2 b0, b1, b2, b3;
decode_fp4x8_e2m1_half4(a_packed, a0, a1, a2, a3);
decode_fp4x8_e2m1_half4(b_packed, b0, b1, b2, b3);
// Vectorized multiplies
__half2 p0 = __hmul2(a0, b0);
__half2 p1 = __hmul2(a1, b1);
__half2 p2 = __hmul2(a2, b2);
__half2 p3 = __hmul2(a3, b3);
// Convert to float and accumulate
float2 fp0 = __half22float2(p0);
float2 fp1 = __half22float2(p1);
float2 fp2 = __half22float2(p2);
float2 fp3 = __half22float2(p3);
block_sum += fp0.x + fp0.y;
block_sum += fp1.x + fp1.y;
block_sum += fp2.x + fp2.y;
block_sum += fp3.x + fp3.y;
}
local_sum += __half2float(scale) * block_sum;
}
}
// Warp-level reduction across K dimension (16 threads per row)
constexpr unsigned int FULL_MASK = 0xffff; // Mask for 16 threads
#pragma unroll
for (int offset = 8; offset > 0; offset >>= 1) {
local_sum += __shfl_down_sync(FULL_MASK, local_sum, offset, 16);
}
// First thread in each row writes to shared memory
if (tidx == 0) {
smem_results[tidy] = __float2half(local_sum);
}
__syncthreads();
// Coalesced write
const int linear_tid = tidx + tidy * THREADS_PER_ROW;
if (linear_tid < ROWS_PER_BLOCK) {
const int write_row = block_row + linear_tid;
if (write_row < M) {
C[batch_idx * M + write_row] = smem_results[linear_tid];
}
}
}
torch::Tensor gemv_fp4_fp8_fp16(
torch::Tensor A,
torch::Tensor B,
torch::Tensor SFA,
torch::Tensor SFB,
torch::Tensor C,
int M, int K, int L)
{
TORCH_CHECK(A.is_cuda(), "A must be a CUDA tensor");
TORCH_CHECK(B.is_cuda(), "B must be a CUDA tensor");
TORCH_CHECK(SFA.is_cuda(), "SFA must be a CUDA tensor");
TORCH_CHECK(SFB.is_cuda(), "SFB must be a CUDA tensor");
TORCH_CHECK(C.is_cuda(), "C must be a CUDA tensor");
constexpr int ROWS_PER_BLOCK = 32;
constexpr int THREADS_PER_ROW = 16;
const dim3 grid((M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK, 1, L);
const dim3 block(THREADS_PER_ROW, ROWS_PER_BLOCK, 1);
block_scaled_gemv_fp4_fp8_fp16_vectorized<<<grid, block>>>(
reinterpret_cast<const uint8_t*>(A.data_ptr()),
reinterpret_cast<const uint8_t*>(B.data_ptr()),
reinterpret_cast<const uint8_t*>(SFA.data_ptr()),
reinterpret_cast<const uint8_t*>(SFB.data_ptr()),
reinterpret_cast<__half*>(C.data_ptr()),
M, K, L);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess)
throw std::runtime_error(cudaGetErrorString(err));
return C;
}
"""
gemv_cpp_src = """
#include <torch/extension.h>
torch::Tensor gemv_fp4_fp8_fp16(
torch::Tensor A,
torch::Tensor B,
torch::Tensor SFA,
torch::Tensor SFB,
torch::Tensor C,
int M, int K, int L);
"""
_gemm_module = load_inline(
name="block_scaled_gemv_vectorized",
cpp_sources=gemv_cpp_src,
cuda_sources=gemv_cuda_src,
functions=["gemv_fp4_fp8_fp16"],
extra_cuda_cflags=[
'-O3',
'--use_fast_math',
'-std=c++17',
'--expt-relaxed-constexpr',
'--maxrregcount=256',
'-gencode=arch=compute_100a,code=sm_100a',
],
verbose=True,
)
def custom_kernel(data: input_t) -> output_t:
"""
K-parallel with explicit vectorized loads and warp-level reduction:
- Keeps original K-parallelism (one thread = multiple tiles)
- Uses __ldg and float4 for explicit vectorized loads
- Warp shuffle reduction
- Vectorized decodes using PTX asm for FP4 x8 and FP8 x4
"""
a, b, sfa, sfb, _, _, c = data
m, k_packed, l = a.shape
k = k_packed * 2
a_uint8 = a.view(torch.uint8)
b_uint8 = b.view(torch.uint8)
sfa_uint8 = sfa.view(torch.uint8)
sfb_uint8 = sfb.view(torch.uint8)
_gemm_module.gemv_fp4_fp8_fp16(a_uint8, b_uint8, sfa_uint8, sfb_uint8, c, m, k, l)
return cscrolls · 283 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 81004.
- # k: 16384; l: 1; m: 7168; seed: 1111- # ⏱ 53.4 ± 0.05 µs- # ⚡ 53.1 µs 🐌 55.3 µs+ import torch+ import sys+ from torch.utils.cpp_extension import load_inline+ from typing import Tuple+ from task import input_t, output_t- # k: 7168; l: 8; m: 4096; seed: 1111- # ⏱ 56.7 ± 0.07 µs- # ⚡ 55.1 µs 🐌 58.4 µs+ gemv_cuda_src = """+ #include <cuda_fp16.h>+ #include <cuda_runtime.h>+ #include <cuda_pipeline.h>+ #include <cstdint>- # k: 2048; l: 4; m: 7168; seed: 1111- # ⏱ 21.1 ± 0.07 µs- # ⚡ 20.4 µs 🐌 22.7 µs- # Params:- # When I use:- # threads_per_m = 128- # # Make sure threads_per_m is divisible by 1024- # threads_per_k = 1024 // threads_per_m- # mma_tiler_mnk = (threads_per_m, 1, 64)- # It's optimal for k=7168:- # k: 16384; l: 1; m: 7168; seed: 1111- # ⏱ 53.4 ± 0.05 µs- # ⚡ 53.1 µs 🐌 55.3 µs+ // Vectorized decode for FP4 x8 - FIXED VERSION+ // cvt.rn.f16x2.e2m1x2 takes .b8 input and produces .b32 output (f16x2)+ __device__ __forceinline__ void decode_fp4x8_e2m1_half4(const uint32_t packed, __half2& h0, __half2& h1, __half2& h2, __half2& h3)+ {+ uint32_t out0, out1, out2, out3;++ // PTX allows wider operands for .b8 instruction types+ // Input is uint32 containing 4 bytes, output is 4x .b32 registers+ asm volatile (+ "{"+ " .reg .b8 %%b0, %%b1, %%b2, %%b3;\\n"+ " mov.b32 {%%b0, %%b1, %%b2, %%b3}, %4;\\n"+ " cvt.rn.f16x2.e2m1x2 %0, %%b0;\\n"+ " cvt.rn.f16x2.e2m1x2 %1, %%b1;\\n"+ " cvt.rn.f16x2.e2m1x2 %2, %%b2;\\n"+ " cvt.rn.f16x2.e2m1x2 %3, %%b3;\\n"+ "}"+ : "=r"(out0), "=r"(out1), "=r"(out2), "=r"(out3)+ : "r"(packed)+ );++ h0 = *reinterpret_cast<const __half2*>(&out0);+ h1 = *reinterpret_cast<const __half2*>(&out1);+ h2 = *reinterpret_cast<const __half2*>(&out2);+ h3 = *reinterpret_cast<const __half2*>(&out3);+ }- # k: 7168; l: 8; m: 4096; seed: 1111- # ⏱ 57.3 ± 0.06 µs- # ⚡ 55.3 µs 🐌 58.4 µs+ // Vectorized decode for FP8 x4 - FIXED VERSION+ // cvt.rn.f16x2.e4m3x2 takes .b16 input and produces .b32 output (f16x2)+ __device__ __forceinline__ void decode_fp8x4_e4m3fn_half4(const uint32_t packed, __half& h0, __half& h1, __half& h2, __half& h3)+ {+ uint32_t out_low, out_high;++ // Input is uint32 containing 2x16-bit values+ asm volatile (+ "{"+ " .reg .b16 %%low, %%high;\\n"+ " mov.b32 {%%low, %%high}, %2;\\n"+ " cvt.rn.f16x2.e4m3x2 %0, %%low;\\n"+ " cvt.rn.f16x2.e4m3x2 %1, %%high;\\n"+ "}"+ : "=r"(out_low), "=r"(out_high)+ : "r"(packed)+ );++ __half2 h_low = *reinterpret_cast<const __half2*>(&out_low);+ __half2 h_high = *reinterpret_cast<const __half2*>(&out_high);+ h0 = h_low.x;+ h1 = h_low.y;+ h2 = h_high.x;+ h3 = h_high.y;+ }- # k: 2048; l: 4; m: 7168; seed: 1111- # ⏱ 22.2 ± 0.05 µs- # ⚡ 20.4 µs 🐌 22.8 µs+ extern "C" __global__+ void block_scaled_gemv_fp4_fp8_fp16_vectorized(+ const uint8_t* __restrict__ A, // [l, m, k//2]+ const uint8_t* __restrict__ B, // [l, 128, k//2]+ const uint8_t* __restrict__ SFA, // [l, m, k//16]+ const uint8_t* __restrict__ SFB, // [l, 128, k//16]+ __half* __restrict__ C, // [l, m, 1]+ int M, int K, int L)+ {+ constexpr int ROWS_PER_BLOCK = 32;+ constexpr int THREADS_PER_ROW = 16; // 16 threads collaborate on each row+ constexpr int TILE_K = 64; // Each thread handles 64 elements+ constexpr int SF_VEC_SIZE = 16;- # When I use:- # threads_per_m = 64- # # Make sure threads_per_m is divisible by 1024- # threads_per_k = 1024 // threads_per_m- # mma_tiler_mnk = (threads_per_m, 1, 64)- # It's optimal for k=16384:- # k: 16384; l: 1; m: 7168; seed: 1111- # ⏱ 34.1 ± 0.06 µs- # ⚡ 32.7 µs 🐌 34.9 µs+ const int tidx = threadIdx.x; // 0-15: which K-segment for this row+ const int tidy = threadIdx.y; // 0-31: which row+ const int block_row = blockIdx.x * ROWS_PER_BLOCK;+ const int batch_idx = blockIdx.z;+ const int global_row = block_row + tidy;- # k: 7168; l: 8; m: 4096; seed: 1111- # ⏱ 57.5 ± 0.06 µs- # ⚡ 56.4 µs 🐌 58.4 µs+ // Shared memory for coalesced writes+ __shared__ __half smem_results[ROWS_PER_BLOCK];- # k: 2048; l: 4; m: 7168; seed: 1111- # ⏱ 24.6 ± 0.02 µs- # ⚡ 23.5 µs 🐌 25.6 µs+ if (global_row >= M) return;- # Kernel configuration parameters- # threads_per_m = 32- # # Make sure threads_per_m is divisible by 512- # threads_per_k = 512 // threads_per_m- # Gives k=16384 best performance.- # k: 16384; l: 1; m: 7168; seed: 1111- # ⏱ 33.8 ± 0.07 µs- # ⚡ 32.7 µs 🐌 35.0 µs+ float local_sum = 0.0f;+ const int num_k_tiles = K / TILE_K;- # k: 7168; l: 8; m: 4096; seed: 1111- # ⏱ 55.3 ± 0.03 µs- # ⚡ 54.2 µs 🐌 56.3 µs+ const uint8_t* B_base = B + batch_idx * (128 * K / 2);+ const uint8_t* SFB_base = SFB + batch_idx * (128 * K / 16);+ const uint8_t* A_row = A + batch_idx * (M * K / 2) + global_row * (K / 2);+ const uint8_t* SFA_row = SFA + batch_idx * (M * K / 16) + global_row * (K / 16);- # k: 2048; l: 4; m: 7168; seed: 1111- # ⏱ 20.6 ± 0.02 µs- # ⚡ 20.4 µs 🐌 22.6 µs+ // K-PARALLELISM: Each thread handles multiple K-tiles, striding by THREADS_PER_ROW+ #pragma unroll 1+ for (int tile = tidx; tile < num_k_tiles; tile += THREADS_PER_ROW) {+ const int k_offset = tile * TILE_K;- import torch- from task import input_t, output_t+ // Vectorized loads+ float4 A_data[2]; // 2x16 bytes = 32 bytes+ float4 B_data[2];- import cutlass- import cutlass.cute as cute- from cutlass.cute.runtime import make_ptr- import cutlass.utils.blockscaled_layout as blockscaled_utils- from cutlass.utils import SmemAllocator+ const float4* A_ptr = reinterpret_cast<const float4*>(A_row + k_offset / 2);+ const float4* B_ptr = reinterpret_cast<const float4*>(B_base + k_offset / 2);- # Kernel configuration parameters- threads_per_m = 32- # Make sure threads_per_m is divisible by 512- threads_per_k = 512 // threads_per_m- mma_tiler_mnk = (threads_per_m, 1, 64)- ab_dtype = cutlass.Float4E2M1FN # FP4 data type for A and B- sf_dtype = cutlass.Float8E4M3FN # FP8 data type for scale factors- c_dtype = cutlass.Float16 # FP16 output type- sf_vec_size = 16 # Scale factor block size (16 elements share one scale)+ A_data[0] = __ldg(A_ptr);+ A_data[1] = __ldg(A_ptr + 1);+ B_data[0] = __ldg(B_ptr);+ B_data[1] = __ldg(B_ptr + 1);- # Helper function for ceiling division- def ceil_div(a, b):- return (a + b - 1) // b+ uint8_t* A_tile_data = reinterpret_cast<uint8_t*>(A_data);+ uint8_t* B_tile_data = reinterpret_cast<uint8_t*>(B_data);- # The CuTe reference implementation for NVFP4 block-scaled GEMV- @cute.kernel- def kernel(- mA_mkl: cute.Tensor,- mB_nkl: cute.Tensor,- mSFA_mkl: cute.Tensor,- mSFB_nkl: cute.Tensor,- mC_mnl: cute.Tensor,- ):- # Get CUDA block and thread indices- bidx, bidy, bidz = cute.arch.block_idx()- tidx, tidy, _ = cute.arch.thread_idx()+ // Load scale factors (4 bytes per thread)+ uint32_t sfa_vec = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + k_offset / 16));+ uint32_t sfb_vec = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + k_offset / 16));- # Extract the local tile for input matrix A (shape: [block_M, block_K, rest_M, rest_K, rest_L])- gA_mkl = cute.local_tile(- mA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)- )- # Extract the local tile for scale factor tensor for A (same shape as gA_mkl)- # Here, block_M = (32, 4); block_K = (16, 4)- gSFA_mkl = cute.local_tile(- mSFA_mkl, cute.slice_(mma_tiler_mnk, (None, 0, None)), (None, None, None)- )- # Extract the local tile for input matrix B (shape: [block_N, block_K, rest_N, rest_K, rest_L])- gB_nkl = cute.local_tile(- mB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)- )- # Extract the local tile for scale factor tensor for B (same shape as gB_nkl)- gSFB_nkl = cute.local_tile(- mSFB_nkl, cute.slice_(mma_tiler_mnk, (0, None, None)), (None, None, None)- )- # Extract the local tile for output matrix C (shape: [block_M, block_N, rest_M, rest_N, rest_L])- gC_mnl = cute.local_tile(- mC_mnl, cute.slice_(mma_tiler_mnk, (None, None, 0)), (None, None, None)- )+ // Vectorized decode for all 4 scales at once+ __half sfa_scales[4], sfb_scales[4];+ decode_fp8x4_e4m3fn_half4(sfa_vec, sfa_scales[0], sfa_scales[1], sfa_scales[2], sfa_scales[3]);+ decode_fp8x4_e4m3fn_half4(sfb_vec, sfb_scales[0], sfb_scales[1], sfb_scales[2], sfb_scales[3]);- # Select output element corresponding to this thread and block indices- tCgC = gC_mnl[tidx, None, bidx, bidy, bidz]- tCgC = cute.make_tensor(tCgC.iterator, 1)- res = cute.zeros_like(tCgC, cutlass.Float32)-- allocator = SmemAllocator()- # Allocate a buffer for row sum accumulation in shared memory- # FIXED: Use stride (1, threads_per_m) to avoid bank conflicts- # With unit stride in tidx dimension, consecutive threads access consecutive addresses- # This ensures conflict-free access when threads write their results- row_sum_buffer = allocator.allocate_tensor(- element_type=cutlass.Float32,- layout=cute.make_layout((threads_per_m, threads_per_k), stride=(1, threads_per_m))- )-- k_tile_cnt = gA_mkl.layout[3].shape- for k_tile in range(tidy, k_tile_cnt, threads_per_k):- tAgA = gA_mkl[tidx, None, bidx, k_tile, bidz]- tBgB = gB_nkl[0, None, bidy, k_tile, bidz]- tAgSFA = gSFA_mkl[tidx, (0, None), bidx, k_tile, bidz]- tBgSFB = gSFB_nkl[0, (0, None), bidy, k_tile, bidz]-- tArA = cute.make_rmem_tensor_like(tAgA, cutlass.Float16)- tBrB = cute.make_rmem_tensor_like(tBgB, cutlass.Float16)- tArSFA = cute.make_rmem_tensor_like(tAgSFA, cutlass.Float32)- tBrSFB = cute.make_rmem_tensor_like(tBgSFB, cutlass.Float32)+ // COMPUTE with vectorized decode using PTX asm+ constexpr int NUM_SF_BLOCKS = TILE_K / SF_VEC_SIZE;+ #pragma unroll+ for (int sf_block = 0; sf_block < NUM_SF_BLOCKS; sf_block++) {+ const int base_k = sf_block * SF_VEC_SIZE;- # Load NVFP4 or FP8 values from global memory- a_val_nvfp4 = tAgA.load()- b_val_nvfp4 = tBgB.load()- sfa_val_fp8 = tAgSFA.load()- sfb_val_fp8 = tBgSFB.load()+ // Use pre-decoded scales+ __half sfa = sfa_scales[sf_block];+ __half sfb = sfb_scales[sf_block];+ __half scale = __hmul(sfa, sfb);+ float block_sum = 0.0f;- # Store the converted values to RMEM CuTe tensors- tArA.store(a_val_nvfp4.to(cutlass.Float16))- tBrB.store(b_val_nvfp4.to(cutlass.Float16))- tArSFA.store(sfa_val_fp8.to(cutlass.Float32))- tBrSFB.store(sfb_val_fp8.to(cutlass.Float32))+ // Process 16 elements (8 bytes packed) using two x8 decodes+ #pragma unroll+ for (int sub = 0; sub < 2; sub++) {+ const int byte_idx = base_k / 2 + sub * 4;- # Iterate over SF vector tiles and compute the scale&matmul accumulation- for sf_block in cutlass.range_constexpr(mma_tiler_mnk[2] // sf_vec_size):- tmp = cute.zeros_like(tCgC, cutlass.Float32)- base = sf_block * sf_vec_size+ const uint32_t a_packed = *reinterpret_cast<const uint32_t*>(A_tile_data + byte_idx);+ const uint32_t b_packed = *reinterpret_cast<const uint32_t*>(B_tile_data + byte_idx);- for offset in cutlass.range_constexpr(sf_vec_size):- tmp += tArA[base + offset] * tBrB[base + offset]- res += tArSFA[sf_block] * tBrSFB[sf_block] * tmp+ __half2 a0, a1, a2, a3;+ __half2 b0, b1, b2, b3;+ decode_fp4x8_e2m1_half4(a_packed, a0, a1, a2, a3);+ decode_fp4x8_e2m1_half4(b_packed, b0, b1, b2, b3);- row_sum_buffer[(tidx, tidy)] = res[0]- cute.arch.sync_threads()-- if tidy == 0:- out = cute.zeros_like(tCgC, cutlass.Float32)- for i in cutlass.range_constexpr(threads_per_k):- out += row_sum_buffer[(tidx, i)]+ // Vectorized multiplies+ __half2 p0 = __hmul2(a0, b0);+ __half2 p1 = __hmul2(a1, b1);+ __half2 p2 = __hmul2(a2, b2);+ __half2 p3 = __hmul2(a3, b3);- # Store the final float16 result back to global memory- tCgC.store(out.to(cutlass.Float16))- return+ // Convert to float and accumulate+ float2 fp0 = __half22float2(p0);+ float2 fp1 = __half22float2(p1);+ float2 fp2 = __half22float2(p2);+ float2 fp3 = __half22float2(p3);- @cute.jit- def my_kernel(- a_ptr: cute.Pointer,- b_ptr: cute.Pointer,- sfa_ptr: cute.Pointer,- sfb_ptr: cute.Pointer,- c_ptr: cute.Pointer,- problem_size: tuple,- ):- """- Host-side JIT function to prepare tensors and launch GPU kernel.- """- m, _, k, l = problem_size- # Create CuTe Tensor via pointer and problem size.- a_tensor = cute.make_tensor(- a_ptr,- cute.make_layout(- (m, cute.assume(k, 32), l),- stride=(cute.assume(k, 32), 1, cute.assume(m * k, 32)),- ),- )- # We use n=128 to create the torch tensor to do fp4 computation via torch._scaled_mm- # then copy torch tensor to cute tensor for cute customize kernel computation- # therefore we need to ensure b_tensor has the right stride with this 128 padded size on n.- n_padded_128 = 128- b_tensor = cute.make_tensor(- b_ptr,- cute.make_layout(- (n_padded_128, cute.assume(k, 32), l),- stride=(cute.assume(k, 32), 1, cute.assume(n_padded_128 * k, 32)),- ),- )- c_tensor = cute.make_tensor(- c_ptr, cute.make_layout((cute.assume(m, 32), 1, l), stride=(1, 1, m))- )- # Convert scale factor tensors to MMA layout- # The layout matches Tensor Core requirements: (((32, 4), REST_M), ((SF_K, 4), REST_K), (1, REST_L))- sfa_layout = blockscaled_utils.tile_atom_to_shape_SF(a_tensor.shape, sf_vec_size)- sfa_tensor = cute.make_tensor(sfa_ptr, sfa_layout)+ block_sum += fp0.x + fp0.y;+ block_sum += fp1.x + fp1.y;+ block_sum += fp2.x + fp2.y;+ block_sum += fp3.x + fp3.y;+ }- sfb_layout = blockscaled_utils.tile_atom_to_shape_SF(b_tensor.shape, sf_vec_size)- sfb_tensor = cute.make_tensor(sfb_ptr, sfb_layout)+ local_sum += __half2float(scale) * block_sum;+ }+ }- # Compute grid dimensions- # Grid is (M_blocks, 1, L) where:- # - M_blocks = ceil(M / 128) to cover all output rows- # - L = batch size- grid = (- cute.ceil_div(c_tensor.shape[0], threads_per_m),- 1,- c_tensor.shape[2],- )+ // Warp-level reduction across K dimension (16 threads per row)+ constexpr unsigned int FULL_MASK = 0xffff; // Mask for 16 threads++ #pragma unroll+ for (int offset = 8; offset > 0; offset >>= 1) {+ local_sum += __shfl_down_sync(FULL_MASK, local_sum, offset, 16);+ }- # Launch the CUDA kernel- kernel(a_tensor, b_tensor, sfa_tensor, sfb_tensor, c_tensor).launch(- grid=grid,- block=[threads_per_m, threads_per_k, 1],- cluster=(1, 1, 1),- )- return+ // First thread in each row writes to shared memory+ if (tidx == 0) {+ smem_results[tidy] = __float2half(local_sum);+ }+ __syncthreads();+ // Coalesced write+ const int linear_tid = tidx + tidy * THREADS_PER_ROW;+ if (linear_tid < ROWS_PER_BLOCK) {+ const int write_row = block_row + linear_tid;+ if (write_row < M) {+ C[batch_idx * M + write_row] = smem_results[linear_tid];+ }+ }+ }- # Global cache for compiled kernel- _compiled_kernel_cache = None+ torch::Tensor gemv_fp4_fp8_fp16(+ torch::Tensor A,+ torch::Tensor B,+ torch::Tensor SFA,+ torch::Tensor SFB,+ torch::Tensor C,+ int M, int K, int L)+ {+ TORCH_CHECK(A.is_cuda(), "A must be a CUDA tensor");+ TORCH_CHECK(B.is_cuda(), "B must be a CUDA tensor");+ TORCH_CHECK(SFA.is_cuda(), "SFA must be a CUDA tensor");+ TORCH_CHECK(SFB.is_cuda(), "SFB must be a CUDA tensor");+ TORCH_CHECK(C.is_cuda(), "C must be a CUDA tensor");+ constexpr int ROWS_PER_BLOCK = 32;+ constexpr int THREADS_PER_ROW = 16;- # This function is used to compile the kernel once and cache it and then allow users to- # run the kernel multiple times to get more accurate timing results.- def compile_kernel():- """- Compile the kernel once and cache it.- This should be called before any timing measurements.+ const dim3 grid((M + ROWS_PER_BLOCK - 1) / ROWS_PER_BLOCK, 1, L);+ const dim3 block(THREADS_PER_ROW, ROWS_PER_BLOCK, 1);- Returns:- The compiled kernel function- """- global _compiled_kernel_cache+ block_scaled_gemv_fp4_fp8_fp16_vectorized<<<grid, block>>>(+ reinterpret_cast<const uint8_t*>(A.data_ptr()),+ reinterpret_cast<const uint8_t*>(B.data_ptr()),+ reinterpret_cast<const uint8_t*>(SFA.data_ptr()),+ reinterpret_cast<const uint8_t*>(SFB.data_ptr()),+ reinterpret_cast<__half*>(C.data_ptr()),+ M, K, L);- if _compiled_kernel_cache is not None:- return _compiled_kernel_cache+ cudaError_t err = cudaGetLastError();+ if (err != cudaSuccess)+ throw std::runtime_error(cudaGetErrorString(err));- # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer- a_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)- b_ptr = make_ptr(ab_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)- c_ptr = make_ptr(c_dtype, 0, cute.AddressSpace.gmem, assumed_align=16)- sfa_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)- sfb_ptr = make_ptr(sf_dtype, 0, cute.AddressSpace.gmem, assumed_align=32)+ return C;+ }+ """- # Compile the kernel- _compiled_kernel_cache = cute.compile(- my_kernel, a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (0, 0, 0, 0)- )+ gemv_cpp_src = """+ #include <torch/extension.h>- return _compiled_kernel_cache+ torch::Tensor gemv_fp4_fp8_fp16(+ torch::Tensor A,+ torch::Tensor B,+ torch::Tensor SFA,+ torch::Tensor SFB,+ torch::Tensor C,+ int M, int K, int L);+ """+ _gemm_module = load_inline(+ name="block_scaled_gemv_vectorized",+ cpp_sources=gemv_cpp_src,+ cuda_sources=gemv_cuda_src,+ functions=["gemv_fp4_fp8_fp16"],+ extra_cuda_cflags=[+ '-O3',+ '--use_fast_math',+ '-std=c++17',+ '--expt-relaxed-constexpr',+ '--maxrregcount=256',+ '-gencode=arch=compute_100a,code=sm_100a',+ ],+ verbose=True,+ )def custom_kernel(data: input_t) -> output_t:"""- Execute the block-scaled GEMV kernel.-- This is the main entry point called by the evaluation framework.- It converts PyTorch tensors to CuTe tensors, launches the kernel,- and returns the result.-- Args:- data: Tuple of (a, b, sfa_cpu, sfb_cpu, c) PyTorch tensors- a: [m, k, l] - Input matrix in float4e2m1fn- b: [1, k, l] - Input vector in float4e2m1fn- sfa_cpu: [m, k, l] - Scale factors in float8_e4m3fn- sfb_cpu: [1, k, l] - Scale factors in float8_e4m3fn- sfa_permuted: [32, 4, rest_m, 4, rest_k, l] - Scale factors in float8_e4m3fn- sfb_permuted: [32, 4, rest_n, 4, rest_k, l] - Scale factors in float8_e4m3fn- c: [m, 1, l] - Output vector in float16-- Returns:- Output tensor c with computed GEMV results+ K-parallel with explicit vectorized loads and warp-level reduction:+ - Keeps original K-parallelism (one thread = multiple tiles)+ - Uses __ldg and float4 for explicit vectorized loads+ - Warp shuffle reduction+ - Vectorized decodes using PTX asm for FP4 x8 and FP8 x4"""- a, b, _, _, sfa_permuted, sfb_permuted, c = data+ a, b, sfa, sfb, _, _, c = data+ m, k_packed, l = a.shape+ k = k_packed * 2- # Ensure kernel is compiled (will use cached version if available)- # To avoid the compilation overhead, we compile the kernel once and cache it.- compiled_func = compile_kernel()+ a_uint8 = a.view(torch.uint8)+ b_uint8 = b.view(torch.uint8)+ sfa_uint8 = sfa.view(torch.uint8)+ sfb_uint8 = sfb.view(torch.uint8)- # Get dimensions from MxKxL layout- m, k, l = a.shape- # Torch use e2m1_x2 data type, thus k is halved- k = k * 2- # GEMV N dimension is always 1- n = 1+ _gemm_module.gemv_fp4_fp8_fp16(a_uint8, b_uint8, sfa_uint8, sfb_uint8, c, m, k, l)- # Create CuTe pointers for A/B/C/SFA/SFB via torch tensor data pointer- a_ptr = make_ptr(ab_dtype, a.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- b_ptr = make_ptr(ab_dtype, b.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)- c_ptr = make_ptr(c_dtype, c.data_ptr(), cute.AddressSpace.gmem, assumed_align=16)-- sfa_ptr = make_ptr(- sf_dtype, sfa_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32- )- sfb_ptr = make_ptr(- sf_dtype, sfb_permuted.data_ptr(), cute.AddressSpace.gmem, assumed_align=32- )-- # Execute the compiled kernel- compiled_func(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, (m, n, k, l))-return cNo newline at end of file
scrolls · 568 diff lines total
Best evidence level for this revision: reported
JSON