submission 105562
yue · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 333 lines, June 9 Researcher Reciprocity License v1.0.
submit_v3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-105562?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:5aaf2135fa7a4b698907994136e8aa7e535285fce0e2bd3234e8ac6b8de4b40c
license declaredunknown
license concludedunknown
authorsyue
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
tile-k = 64
constexpr int TILE_K = 64;vector-width = float4
const float4* A_ptr_0 = reinterpret_cast<const float4*>(A_row + (k_offset_0 >> 1));Kernel source
submit_v3.py333 lines
import torch
import sys
from torch.utils.cpp_extension import load_inline
from typing import Tuple
gemv_cuda_src = """
#include <cuda_fp16.h>
#include <cuda_runtime.h>
#include <cstdint>
__device__ __forceinline__ float decode_mul_accumulate_fp4x8(
const uint32_t a_packed,
const uint32_t b_packed,
float acc)
{
float result;
asm volatile (
"{"
" .reg .b8 %%ab<4>, %%bb<4>;\\n"
" .reg .b32 %%a<4>, %%b<4>;\\n"
" .reg .b32 %%p0, %%p1;\\n"
" .reg .f16 %%h0, %%h1;\\n"
" .reg .f32 %%f0, %%f1;\\n"
" mov.b32 {%%ab0, %%ab1, %%ab2, %%ab3}, %1;\\n"
" mov.b32 {%%bb0, %%bb1, %%bb2, %%bb3}, %2;\\n"
" cvt.rn.f16x2.e2m1x2 %%a0, %%ab0;\\n"
" cvt.rn.f16x2.e2m1x2 %%a1, %%ab1;\\n"
" cvt.rn.f16x2.e2m1x2 %%a2, %%ab2;\\n"
" cvt.rn.f16x2.e2m1x2 %%a3, %%ab3;\\n"
" cvt.rn.f16x2.e2m1x2 %%b0, %%bb0;\\n"
" cvt.rn.f16x2.e2m1x2 %%b1, %%bb1;\\n"
" cvt.rn.f16x2.e2m1x2 %%b2, %%bb2;\\n"
" cvt.rn.f16x2.e2m1x2 %%b3, %%bb3;\\n"
" mul.rn.f16x2 %%p0, %%a0, %%b0;\\n"
" fma.rn.f16x2 %%p0, %%a1, %%b1, %%p0;\\n"
" mul.rn.f16x2 %%p1, %%a2, %%b2;\\n"
" fma.rn.f16x2 %%p1, %%a3, %%b3, %%p1;\\n"
" add.rn.f16x2 %%p0, %%p0, %%p1;\\n"
" mov.b32 {%%h0, %%h1}, %%p0;\\n"
" cvt.f32.f16 %%f0, %%h0;\\n"
" cvt.f32.f16 %%f1, %%h1;\\n"
" add.f32 %%f0, %%f0, %%f1;\\n"
" add.f32 %0, %%f0, %3;\\n"
"}"
: "=f"(result)
: "r"(a_packed), "r"(b_packed), "f"(acc)
);
return result;
}
__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;
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;
}
// Process one tile and return partial sum
__device__ __forceinline__ float process_tile(
const uint32_t* A_u32_0,
const uint32_t* B_u32_0,
const uint32_t* A_u32_1,
const uint32_t* B_u32_1,
const __half* sfa_scales,
const __half* sfb_scales)
{
float tile_sum = 0.0f;
// SF block 0
{
__half scale = __hmul(sfa_scales[0], sfb_scales[0]);
float block_sum = 0.0f;
block_sum = decode_mul_accumulate_fp4x8(A_u32_0[0], B_u32_0[0], block_sum);
block_sum = decode_mul_accumulate_fp4x8(A_u32_0[1], B_u32_0[1], block_sum);
tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
}
// SF block 1
{
__half scale = __hmul(sfa_scales[1], sfb_scales[1]);
float block_sum = 0.0f;
block_sum = decode_mul_accumulate_fp4x8(A_u32_0[2], B_u32_0[2], block_sum);
block_sum = decode_mul_accumulate_fp4x8(A_u32_0[3], B_u32_0[3], block_sum);
tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
}
// SF block 2
{
__half scale = __hmul(sfa_scales[2], sfb_scales[2]);
float block_sum = 0.0f;
block_sum = decode_mul_accumulate_fp4x8(A_u32_1[0], B_u32_1[0], block_sum);
block_sum = decode_mul_accumulate_fp4x8(A_u32_1[1], B_u32_1[1], block_sum);
tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
}
// SF block 3
{
__half scale = __hmul(sfa_scales[3], sfb_scales[3]);
float block_sum = 0.0f;
block_sum = decode_mul_accumulate_fp4x8(A_u32_1[2], B_u32_1[2], block_sum);
block_sum = decode_mul_accumulate_fp4x8(A_u32_1[3], B_u32_1[3], block_sum);
tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);
}
return tile_sum;
}
extern "C" __global__ __launch_bounds__(128, 8)
void block_scaled_gemv_fp4_fp8_fp16_optimized(
const uint8_t* __restrict__ A,
const uint8_t* __restrict__ B,
const uint8_t* __restrict__ SFA,
const uint8_t* __restrict__ SFB,
__half* __restrict__ C,
int M, int K, int L)
{
constexpr int ROWS_PER_BLOCK = 8;
constexpr int THREADS_PER_ROW = 16;
constexpr int TILE_K = 64;
const int tidx = threadIdx.x;
const int tidy = threadIdx.y;
const int block_row = blockIdx.x * ROWS_PER_BLOCK;
const int batch_idx = blockIdx.z;
const int global_row = block_row + tidy;
if (global_row >= M) return;
const int num_k_tiles = K >> 6;
const uint8_t* B_base = B + batch_idx * (128 * (K >> 1));
const uint8_t* SFB_base = SFB + batch_idx * (128 * (K >> 4));
const uint8_t* A_row = A + batch_idx * (M * (K >> 1)) + global_row * (K >> 1);
const uint8_t* SFA_row = SFA + batch_idx * (M * (K >> 4)) + global_row * (K >> 4);
float local_sum = 0.0f;
// Main loop - process 2 tiles per iteration for better ILP
int tile = tidx;
#pragma unroll
for (; tile + THREADS_PER_ROW < num_k_tiles; tile += 2 * THREADS_PER_ROW) {
// Load tile 0
const int k_offset_0 = tile * TILE_K;
const float4* A_ptr_0 = reinterpret_cast<const float4*>(A_row + (k_offset_0 >> 1));
const float4* B_ptr_0 = reinterpret_cast<const float4*>(B_base + (k_offset_0 >> 1));
float4 A_data0_t0 = __ldg(A_ptr_0);
float4 B_data0_t0 = __ldg(B_ptr_0);
float4 A_data1_t0 = __ldg(A_ptr_0 + 1);
float4 B_data1_t0 = __ldg(B_ptr_0 + 1);
uint32_t sfa_vec_t0 = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset_0 >> 4)));
uint32_t sfb_vec_t0 = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset_0 >> 4)));
// Load tile 1
const int k_offset_1 = (tile + THREADS_PER_ROW) * TILE_K;
const float4* A_ptr_1 = reinterpret_cast<const float4*>(A_row + (k_offset_1 >> 1));
const float4* B_ptr_1 = reinterpret_cast<const float4*>(B_base + (k_offset_1 >> 1));
float4 A_data0_t1 = __ldg(A_ptr_1);
float4 B_data0_t1 = __ldg(B_ptr_1);
float4 A_data1_t1 = __ldg(A_ptr_1 + 1);
float4 B_data1_t1 = __ldg(B_ptr_1 + 1);
uint32_t sfa_vec_t1 = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset_1 >> 4)));
uint32_t sfb_vec_t1 = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset_1 >> 4)));
// Process tile 0
__half sfa_scales_t0[4], sfb_scales_t0[4];
decode_fp8x4_e4m3fn_half4(sfa_vec_t0, sfa_scales_t0[0], sfa_scales_t0[1], sfa_scales_t0[2], sfa_scales_t0[3]);
decode_fp8x4_e4m3fn_half4(sfb_vec_t0, sfb_scales_t0[0], sfb_scales_t0[1], sfb_scales_t0[2], sfb_scales_t0[3]);
local_sum += process_tile(
reinterpret_cast<const uint32_t*>(&A_data0_t0),
reinterpret_cast<const uint32_t*>(&B_data0_t0),
reinterpret_cast<const uint32_t*>(&A_data1_t0),
reinterpret_cast<const uint32_t*>(&B_data1_t0),
sfa_scales_t0, sfb_scales_t0);
// Process tile 1
__half sfa_scales_t1[4], sfb_scales_t1[4];
decode_fp8x4_e4m3fn_half4(sfa_vec_t1, sfa_scales_t1[0], sfa_scales_t1[1], sfa_scales_t1[2], sfa_scales_t1[3]);
decode_fp8x4_e4m3fn_half4(sfb_vec_t1, sfb_scales_t1[0], sfb_scales_t1[1], sfb_scales_t1[2], sfb_scales_t1[3]);
local_sum += process_tile(
reinterpret_cast<const uint32_t*>(&A_data0_t1),
reinterpret_cast<const uint32_t*>(&B_data0_t1),
reinterpret_cast<const uint32_t*>(&A_data1_t1),
reinterpret_cast<const uint32_t*>(&B_data1_t1),
sfa_scales_t1, sfb_scales_t1);
}
// Handle remaining tile if odd number
if (tile < num_k_tiles) {
const int k_offset = tile * TILE_K;
const float4* A_ptr = reinterpret_cast<const float4*>(A_row + (k_offset >> 1));
const float4* B_ptr = reinterpret_cast<const float4*>(B_base + (k_offset >> 1));
float4 A_data0 = __ldg(A_ptr);
float4 B_data0 = __ldg(B_ptr);
float4 A_data1 = __ldg(A_ptr + 1);
float4 B_data1 = __ldg(B_ptr + 1);
uint32_t sfa_vec = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset >> 4)));
uint32_t sfb_vec = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset >> 4)));
__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]);
local_sum += process_tile(
reinterpret_cast<const uint32_t*>(&A_data0),
reinterpret_cast<const uint32_t*>(&B_data0),
reinterpret_cast<const uint32_t*>(&A_data1),
reinterpret_cast<const uint32_t*>(&B_data1),
sfa_scales, sfb_scales);
}
// Warp-level reduction
#pragma unroll
for (int offset = 8; offset > 0; offset >>= 1) {
local_sum += __shfl_xor_sync(0xffff, local_sum, offset, 16);
}
if (tidx == 0) {
C[batch_idx * M + global_row] = __float2half(local_sum);
}
}
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 = 8;
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_optimized<<<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_ilp_v1",
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=255',
'--prec-div=false',
'--fmad=true',
'--ftz=true',
'-gencode=arch=compute_100a,code=sm_100a',
],
verbose=True,
)
def custom_kernel(data):
"""
Optimized version focusing on ILP without shared memory:
- Process 2 tiles per iteration to increase ILP
- All loads issued together, then all computes
- __launch_bounds__ for occupancy hint
- No shared memory overhead
- Direct global -> register path via __ldg
"""
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 · 333 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 105222.
⋯ 1 unchanged linesimport sysfrom torch.utils.cpp_extension import load_inlinefrom typing import Tuple- from task import input_t, output_tgemv_cuda_src = """#include <cuda_fp16.h>#include <cuda_runtime.h>- #include <cuda_pipeline.h>#include <cstdint>__device__ __forceinline__ float decode_mul_accumulate_fp4x8(⋯ 61 unchanged linesh3 = 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]+ // Process one tile and return partial sum+ __device__ __forceinline__ float process_tile(+ const uint32_t* A_u32_0,+ const uint32_t* B_u32_0,+ const uint32_t* A_u32_1,+ const uint32_t* B_u32_1,+ const __half* sfa_scales,+ const __half* sfb_scales)+ {+ float tile_sum = 0.0f;++ // SF block 0+ {+ __half scale = __hmul(sfa_scales[0], sfb_scales[0]);+ float block_sum = 0.0f;+ block_sum = decode_mul_accumulate_fp4x8(A_u32_0[0], B_u32_0[0], block_sum);+ block_sum = decode_mul_accumulate_fp4x8(A_u32_0[1], B_u32_0[1], block_sum);+ tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);+ }++ // SF block 1+ {+ __half scale = __hmul(sfa_scales[1], sfb_scales[1]);+ float block_sum = 0.0f;+ block_sum = decode_mul_accumulate_fp4x8(A_u32_0[2], B_u32_0[2], block_sum);+ block_sum = decode_mul_accumulate_fp4x8(A_u32_0[3], B_u32_0[3], block_sum);+ tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);+ }++ // SF block 2+ {+ __half scale = __hmul(sfa_scales[2], sfb_scales[2]);+ float block_sum = 0.0f;+ block_sum = decode_mul_accumulate_fp4x8(A_u32_1[0], B_u32_1[0], block_sum);+ block_sum = decode_mul_accumulate_fp4x8(A_u32_1[1], B_u32_1[1], block_sum);+ tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);+ }++ // SF block 3+ {+ __half scale = __hmul(sfa_scales[3], sfb_scales[3]);+ float block_sum = 0.0f;+ block_sum = decode_mul_accumulate_fp4x8(A_u32_1[2], B_u32_1[2], block_sum);+ block_sum = decode_mul_accumulate_fp4x8(A_u32_1[3], B_u32_1[3], block_sum);+ tile_sum = __fmaf_rn(__half2float(scale), block_sum, tile_sum);+ }++ return tile_sum;+ }++ extern "C" __global__ __launch_bounds__(128, 8)+ void block_scaled_gemv_fp4_fp8_fp16_optimized(+ const uint8_t* __restrict__ A,+ const uint8_t* __restrict__ B,+ const uint8_t* __restrict__ SFA,+ const uint8_t* __restrict__ SFB,+ __half* __restrict__ C,int M, int K, int L){constexpr int ROWS_PER_BLOCK = 8;- 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;+ constexpr int THREADS_PER_ROW = 16;+ constexpr int TILE_K = 64;- const int tidx = threadIdx.x; // 0-15: which K-segment for this row- const int tidy = threadIdx.y; // 0-31: which row+ const int tidx = threadIdx.x;+ const int tidy = threadIdx.y;const int block_row = blockIdx.x * ROWS_PER_BLOCK;const int batch_idx = blockIdx.z;const int global_row = block_row + tidy;if (global_row >= M) return;- float local_sum = 0.0f;- const int num_k_tiles = K / TILE_K;+ const int num_k_tiles = K >> 6;- 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);+ const uint8_t* B_base = B + batch_idx * (128 * (K >> 1));+ const uint8_t* SFB_base = SFB + batch_idx * (128 * (K >> 4));+ const uint8_t* A_row = A + batch_idx * (M * (K >> 1)) + global_row * (K >> 1);+ const uint8_t* SFA_row = SFA + batch_idx * (M * (K >> 4)) + global_row * (K >> 4);- // 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;+ float local_sum = 0.0f;- // Vectorized loads- float4 A_data[2]; // 2x16 bytes = 32 bytes- float4 B_data[2];+ // Main loop - process 2 tiles per iteration for better ILP+ int tile = tidx;++ #pragma unroll+ for (; tile + THREADS_PER_ROW < num_k_tiles; tile += 2 * THREADS_PER_ROW) {+ // Load tile 0+ const int k_offset_0 = tile * TILE_K;+ const float4* A_ptr_0 = reinterpret_cast<const float4*>(A_row + (k_offset_0 >> 1));+ const float4* B_ptr_0 = reinterpret_cast<const float4*>(B_base + (k_offset_0 >> 1));++ float4 A_data0_t0 = __ldg(A_ptr_0);+ float4 B_data0_t0 = __ldg(B_ptr_0);+ float4 A_data1_t0 = __ldg(A_ptr_0 + 1);+ float4 B_data1_t0 = __ldg(B_ptr_0 + 1);+ uint32_t sfa_vec_t0 = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset_0 >> 4)));+ uint32_t sfb_vec_t0 = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset_0 >> 4)));- 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);+ // Load tile 1+ const int k_offset_1 = (tile + THREADS_PER_ROW) * TILE_K;+ const float4* A_ptr_1 = reinterpret_cast<const float4*>(A_row + (k_offset_1 >> 1));+ const float4* B_ptr_1 = reinterpret_cast<const float4*>(B_base + (k_offset_1 >> 1));++ float4 A_data0_t1 = __ldg(A_ptr_1);+ float4 B_data0_t1 = __ldg(B_ptr_1);+ float4 A_data1_t1 = __ldg(A_ptr_1 + 1);+ float4 B_data1_t1 = __ldg(B_ptr_1 + 1);+ uint32_t sfa_vec_t1 = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset_1 >> 4)));+ uint32_t sfb_vec_t1 = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset_1 >> 4)));- 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);+ // Process tile 0+ __half sfa_scales_t0[4], sfb_scales_t0[4];+ decode_fp8x4_e4m3fn_half4(sfa_vec_t0, sfa_scales_t0[0], sfa_scales_t0[1], sfa_scales_t0[2], sfa_scales_t0[3]);+ decode_fp8x4_e4m3fn_half4(sfb_vec_t0, sfb_scales_t0[0], sfb_scales_t0[1], sfb_scales_t0[2], sfb_scales_t0[3]);++ local_sum += process_tile(+ reinterpret_cast<const uint32_t*>(&A_data0_t0),+ reinterpret_cast<const uint32_t*>(&B_data0_t0),+ reinterpret_cast<const uint32_t*>(&A_data1_t0),+ reinterpret_cast<const uint32_t*>(&B_data1_t0),+ sfa_scales_t0, sfb_scales_t0);- uint8_t* A_tile_data = reinterpret_cast<uint8_t*>(A_data);- uint8_t* B_tile_data = reinterpret_cast<uint8_t*>(B_data);+ // Process tile 1+ __half sfa_scales_t1[4], sfb_scales_t1[4];+ decode_fp8x4_e4m3fn_half4(sfa_vec_t1, sfa_scales_t1[0], sfa_scales_t1[1], sfa_scales_t1[2], sfa_scales_t1[3]);+ decode_fp8x4_e4m3fn_half4(sfb_vec_t1, sfb_scales_t1[0], sfb_scales_t1[1], sfb_scales_t1[2], sfb_scales_t1[3]);++ local_sum += process_tile(+ reinterpret_cast<const uint32_t*>(&A_data0_t1),+ reinterpret_cast<const uint32_t*>(&B_data0_t1),+ reinterpret_cast<const uint32_t*>(&A_data1_t1),+ reinterpret_cast<const uint32_t*>(&B_data1_t1),+ sfa_scales_t1, sfb_scales_t1);+ }++ // Handle remaining tile if odd number+ if (tile < num_k_tiles) {+ const int k_offset = tile * TILE_K;+ const float4* A_ptr = reinterpret_cast<const float4*>(A_row + (k_offset >> 1));+ const float4* B_ptr = reinterpret_cast<const float4*>(B_base + (k_offset >> 1));- // 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));+ float4 A_data0 = __ldg(A_ptr);+ float4 B_data0 = __ldg(B_ptr);+ float4 A_data1 = __ldg(A_ptr + 1);+ float4 B_data1 = __ldg(B_ptr + 1);+ uint32_t sfa_vec = __ldg(reinterpret_cast<const uint32_t*>(SFA_row + (k_offset >> 4)));+ uint32_t sfb_vec = __ldg(reinterpret_cast<const uint32_t*>(SFB_base + (k_offset >> 4)));- // 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 optimized decode+FMA- #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);-- // Optimized: decode and accumulate in fewer instructions- block_sum = decode_mul_accumulate_fp4x8(a_packed, b_packed, block_sum);- }-- local_sum = __fmaf_rn(__half2float(scale), block_sum, local_sum);- }+ local_sum += process_tile(+ reinterpret_cast<const uint32_t*>(&A_data0),+ reinterpret_cast<const uint32_t*>(&B_data0),+ reinterpret_cast<const uint32_t*>(&A_data1),+ reinterpret_cast<const uint32_t*>(&B_data1),+ sfa_scales, sfb_scales);}- // Warp-level reduction across K dimension (16 threads per row)- constexpr unsigned int FULL_MASK = 0xffff; // Mask for 16 threads-+ // Warp-level reduction#pragma unrollfor (int offset = 8; offset > 0; offset >>= 1) {- local_sum += __shfl_down_sync(FULL_MASK, local_sum, offset, 16);+ local_sum += __shfl_xor_sync(0xffff, local_sum, offset, 16);}- // First thread in each row writes directly to global memoryif (tidx == 0) {C[batch_idx * M + global_row] = __float2half(local_sum);}⋯ 19 unchanged linesconst 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>>>(+ block_scaled_gemv_fp4_fp8_fp16_optimized<<<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()),⋯ 22 unchanged lines"""_gemm_module = load_inline(- name="block_scaled_gemv_vectorized_v1",+ name="block_scaled_gemv_ilp_v1",cpp_sources=gemv_cpp_src,cuda_sources=gemv_cuda_src,functions=["gemv_fp4_fp8_fp16"],⋯ 2 unchanged lines'--use_fast_math','-std=c++17','--expt-relaxed-constexpr',- '--maxrregcount=64',+ '--maxrregcount=255','--prec-div=false','--fmad=true','--ftz=true',⋯ 2 unchanged linesverbose=True,)- def custom_kernel(data: input_t) -> output_t:+ def custom_kernel(data):"""- Optimized K-parallel with PTX FMA instructions:- - K-parallelism (one thread = multiple tiles)- - Vectorized loads using __ldg and float4- - Warp shuffle reduction- - Optimized PTX decode+FMA for FP4 x8 (reduced instruction count)- - PTX FMA for scale multiplication and accumulation+ Optimized version focusing on ILP without shared memory:+ - Process 2 tiles per iteration to increase ILP+ - All loads issued together, then all computes+ - __launch_bounds__ for occupancy hint+ - No shared memory overhead+ - Direct global -> register path via __ldg"""a, b, sfa, sfb, _, _, c = datam, k_packed, l = a.shape⋯ 6 unchanged lines_gemm_module.gemv_fp4_fp8_fp16(a_uint8, b_uint8, sfa_uint8, sfb_uint8, c, m, k, l)- return c-+ return cNo newline at end of file
scrolls · 311 diff lines total
Best evidence level for this revision: reported
JSON