submission 80454
txacvalh · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 330 lines, June 9 Researcher Reciprocity License v1.0.
good_kernel.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-80454?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:609f6daabbedc7ffb6116f058408b100b940575540a2564a017ff319516f97f0
license declaredunknown
license concludedunknown
authorstxacvalh
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
PyTorch reference implementation of NVFP4 block-scaled GEMV.fp8
using fp8e4m3 = __nv_fp8_e4m3;shared-memory
extern __shared__ __align__(BYTES_PER_THREAD_PER_K_TILE) uint32_t sB_u32[];tile-k = 16
constexpr int TILE_K = 16;tile-m = 16
constexpr int TILE_M = 16;tile-n = 8
constexpr int TILE_N = 8;vector-width = half2
half2 tile_acc = reinterpret_cast<__half2 &>(tile_acc2);Kernel source
good_kernel.py330 lines
import torch
from torch.utils.cpp_extension import load_inline
from typing import List
from task import input_t, output_t
import os
custom_func_cuda_source = """
// m16n8k64
using byte = uint8_t;
using int16 = int16_t;
using int32 = int32_t;
using int64 = int64_t;
using fp32 = float;
using fp16 = half;
using bf16 = nv_bfloat16;
using fp8e4m3 = __nv_fp8_e4m3;
using fp8e5m2 = __nv_fp8_e5m2;
using fp8e4m3x4 = uint32_t;
using fp4e2m1 = __nv_fp4_e2m1;
using fp4e2m1x2 = __nv_fp4x2_e2m1;
using fp4e2m1x4 = __nv_fp4x4_e2m1;
using e8m0_t = uint8_t;
constexpr int WARP_SIZE = 32;
constexpr int WARP_COUNT = 8;
constexpr int THREADS_PER_BLOCK = WARP_SIZE * WARP_COUNT;
constexpr int SM_COUNT = 148;
constexpr int FP4X8_SIZE_IN_BYTES = 4;
constexpr int FP4X16_SIZE_IN_BYTES = 8;
constexpr int BLOCK_M = WARP_COUNT;
// constexpr int BLOCK_M = 128;
constexpr int TILE_M = 16;
constexpr int TILE_K = 16;
constexpr int TILE_N = 8;
constexpr int N_PADDED_128 = 128;
constexpr int N = 1;
constexpr int BYTES_PER_THREAD_PER_K_TILE = 32;
template <typename T>
constexpr T DIVUP(const T &x, const T &y) {
return (((x) + ((y)-1)) / (y));
}
__inline__ __device__
int scale_idx(int mn_idx, int k_idx, int l_idx, int MN, int K) {
constexpr int ATOM_SIZE = 128 * 4;
int TOTAL_NUM_ELEMENTS_PER_L = MN * K / 16;
int rest_k = k_idx / 4;
int sub_k_idx = k_idx % 4;
int rest_mn = mn_idx / 128;
int sub_mn_idx1 = (mn_idx % 128) / 32;
int sub_mn_idx0 = mn_idx % 32;
return l_idx * (TOTAL_NUM_ELEMENTS_PER_L) + rest_mn * (128 * K / 16) + rest_k * ATOM_SIZE + sub_mn_idx0 * 16 + sub_mn_idx1 * 4 + sub_k_idx;
}
union fp4vec {
uint64_t vec;
fp4e2m1x2 small_vec[8];
};
/* __device__ __forceinline__ fp8e4m3 pick_byte(uint32_t x, int byte_idx) {
uint32_t tmp;
// byte_idx in [0, 3], bit offset = byte_idx * 8, width = 8
asm volatile (
"bfe.u32 %0, %1, %2, 8;\n"
: "=r"(tmp)
: "r"(x), "r"(byte_idx * 8)
);
return static_cast<uint8_t>(tmp);
} */
__inline__ __device__
uint32_t dot_fp4x8(const uint32_t a, const uint32_t b, uint32_t acc) {
asm volatile( \\
"{\\n" \\
".reg .b8 byte0, byte1, byte2, byte3;\\n" \\
".reg .b8 byte4, byte5, byte6, byte7;\\n" \\
".reg .f16x2 cvt_0, cvt_1, cvt_2, cvt_3;\\n" \\
".reg .f16x2 cvt_4, cvt_5, cvt_6, cvt_7;\\n" \\
"mov.b32 {byte0, byte1, byte2, byte3}, %1;\\n" \\
"mov.b32 {byte4, byte5, byte6, byte7}, %2;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_0, byte0;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_1, byte1;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_2, byte2;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_3, byte3;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_4, byte4;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_5, byte5;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_6, byte6;\\n" \\
"cvt.rn.f16x2.e2m1x2 cvt_7, byte7;\\n" \\
"fma.rn.f16x2 %0, cvt_0, cvt_4, %0;\\n" \\
"fma.rn.f16x2 %0, cvt_1, cvt_5, %0;\\n" \\
"fma.rn.f16x2 %0, cvt_2, cvt_6, %0;\\n" \\
"fma.rn.f16x2 %0, cvt_3, cvt_7, %0;\\n" \\
"}\\n" : "+r"(acc) : "r"(a) , "r"(b));
return acc;
}
// Warp-level reduction: sums 'val' across the active threads in the warp
__inline__ __device__
float warpReduceSum(float val) {
// Do tree-reduction within the warp
for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
val += __shfl_down_sync(0xffffffff, val, offset);
}
return val; // After this, lane 0 holds the sum of the warp
}
__global__ void __maxnreg__(32) custom_func_kernel(const uint8_t* const __restrict__ A,
const uint8_t* const __restrict__ B,
fp8e4m3 * __restrict__ scale_A,
const uint32_t * const __restrict__ scale_B,
half* const __restrict__ C,
const int L, const int M, const int K, const int sB_size_in_bytes) {
const int block_m_idx = blockIdx.x;
const int l_idx = blockIdx.y;
const int warp_id = threadIdx.x / WARP_SIZE;
const int lane_id = threadIdx.x % WARP_SIZE;
int m_idx = block_m_idx * BLOCK_M + warp_id;
extern __shared__ __align__(BYTES_PER_THREAD_PER_K_TILE) uint32_t sB_u32[];
auto sB_scale = reinterpret_cast<fp8e4m3x4 *>(
reinterpret_cast<uint8_t *>(sB_u32) + sB_size_in_bytes);
// Load B into shared mem
int b_l_base_off = (l_idx * N_PADDED_128 * K) / (2);
int b_bound = K / (2 * BYTES_PER_THREAD_PER_K_TILE);
for (int i = threadIdx.x; i < b_bound; i += THREADS_PER_BLOCK) {
uint32_t rB_values[8];
{
void *ptr = (void *)(&B[b_l_base_off + i*BYTES_PER_THREAD_PER_K_TILE]);
asm volatile( \\
"{\\n" \\
"ld.global.v8.b32 { %0, %1, %2, %3, %4, %5, %6, %7 }, [%8];\\n" \\
"}\\n" : "=r"(rB_values[0]), "=r"(rB_values[1]), "=r"(rB_values[2]), "=r"(rB_values[3]), "=r"(rB_values[4]), "=r"(rB_values[5]), "=r"(rB_values[6]), "=r"(rB_values[7]) : "l"(ptr));
}
int rest_idx = i / WARP_SIZE;
#pragma unroll(8)
for (int j = 0; j < 8; j++) {
sB_u32[rest_idx * 256 + j * 32 + lane_id] = rB_values[j];
}
}
int b_scale_bound = K / (16 * sizeof(uint32_t));
auto curr_scale_B = &scale_B[(l_idx * N_PADDED_128 * K / (16 * sizeof(uint32_t)))];
#pragma unroll
for (int i = threadIdx.x; i < b_scale_bound; i += THREADS_PER_BLOCK) {
reinterpret_cast<uint32_t *>(sB_scale)[i] = curr_scale_B[i];
}
__syncthreads();
for (;m_idx < M; m_idx += gridDim.x * BLOCK_M) {
const int a_l_base_off = (l_idx * (M * K) + m_idx * (K)) / 2;
constexpr int ATOM_SIZE = 128 * 4;
const int TOTAL_NUM_ELEMENTS_PER_L = M * K / 16;
const int rest_m = m_idx / 128;
const int sub_m_idx1 = (m_idx % 128) / 32;
const int sub_m_idx0 = m_idx % 32;
const int scale_A_base = l_idx * (TOTAL_NUM_ELEMENTS_PER_L) + m_idx * (K / 16);
auto curr_scale_A = reinterpret_cast<fp8e4m3x4 * >(&scale_A[scale_A_base]);
half acc{};
for (int tile_idx_k = lane_id; tile_idx_k < K / (2 * BYTES_PER_THREAD_PER_K_TILE); tile_idx_k += WARP_SIZE) {
constexpr int factor = (2 * BYTES_PER_THREAD_PER_K_TILE) / (16 * sizeof(fp8e4m3x4));
uint32_t a_scales = curr_scale_A[tile_idx_k * factor];
uint32_t b_scales = sB_scale[tile_idx_k];
union {
uint32_t u;
fp8e4m3 f[4];
} cvt_a, cvt_b;
cvt_a.u = a_scales;
cvt_b.u = b_scales;
uint32_t rA_values[8];
{
void *ptr = (void *)(&A[a_l_base_off + tile_idx_k*BYTES_PER_THREAD_PER_K_TILE]);
asm volatile( \\
"{\\n" \\
"ld.global.v8.b32 { %0, %1, %2, %3, %4, %5, %6, %7 }, [%8];\\n" \\
"}\\n" : "=r"(rA_values[0]), "=r"(rA_values[1]), "=r"(rA_values[2]), "=r"(rA_values[3]), "=r"(rA_values[4]), "=r"(rA_values[5]), "=r"(rA_values[6]), "=r"(rA_values[7]) : "l"(ptr));
}
const int base_off = (tile_idx_k / WARP_SIZE)*256 + lane_id;
#pragma unroll(4)
for (int i = 0; i < 8; i+=2) {
uint32_t tile_acc2{};
tile_acc2 = dot_fp4x8(rA_values[i], sB_u32[base_off + 32 * i], tile_acc2);
tile_acc2 = dot_fp4x8(rA_values[i+1], sB_u32[base_off + 32 * (i + 1)], tile_acc2);
half2 tile_acc = reinterpret_cast<__half2 &>(tile_acc2);
acc += static_cast<half>(tile_acc.x + tile_acc.y) * static_cast<half>(cvt_a.f[i>>1]) * static_cast<half>(cvt_b.f[i>>1]);
}
}
acc = warpReduceSum(acc);
if (lane_id == 0) {
const int c_l_base_off_half = l_idx * (M) + m_idx;
C[c_l_base_off_half] = static_cast<half>(acc);
}
}
}
void custom_func(torch::Tensor A, torch::Tensor B, torch::Tensor scale_A, torch::Tensor scale_B, torch::Tensor C) {
const int L = A.size(0);
const int M = A.size(1);
const int K = A.size(2) * 2;
const int threads = THREADS_PER_BLOCK;
// const int blocks = (N + threads - 1) / threads;
constexpr int sB_block_size = 256*4;
const int sB_size_in_bytes = cuda::ceil_div(K / 2, sB_block_size) * sB_block_size;
const int sB_scale_size_in_bytes = K / 16;
const int shared_mem_size_in_bytes = sB_size_in_bytes + sB_scale_size_in_bytes;
const int GRID_DIM_X = cuda::ceil_div(M, BLOCK_M);
dim3 blocks(GRID_DIM_X, L, 1);
custom_func_kernel<<<blocks, threads, shared_mem_size_in_bytes>>>(
reinterpret_cast<uint8_t *>(A.data_ptr()),
reinterpret_cast<uint8_t *>(B.data_ptr()),
reinterpret_cast<fp8e4m3 *>(scale_A.data_ptr()),
reinterpret_cast<uint32_t *>(scale_B.data_ptr()),
reinterpret_cast<half *>(C.data_ptr()),
L, M, K, sB_size_in_bytes);
/* cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error(cudaGetErrorString(err));
} */
}
"""
# print(custom_func_cuda_source)
custom_func_cpp_source = """
#include <cuda.h>
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_fp4.h>
#include <torch/extension.h>
void custom_func(torch::Tensor A, torch::Tensor B, torch::Tensor scale_A, torch::Tensor scale_B, torch::Tensor C);
"""
os.environ["TORCH_CUDA_ARCH_LIST"] = "10.0a+PTX;12.0a"
module = load_inline(
name='custom_func',
cpp_sources=custom_func_cpp_source,
cuda_sources=custom_func_cuda_source,
functions=['custom_func'],
verbose=True,
extra_cuda_cflags=['-U__CUDA_NO_HALF_CONVERSIONS__', '-U__CUDA_NO_HALF_OPERATORS__ ', '-U__CUDA_NO_HALF2_OPERATORS__', '-O3', '-Xptxas', '-O3', '-use_fast_math'],
)
def custom_kernel(
data: input_t,
) -> output_t:
"""
Reference implementation of block-scale fp8 gemv
Args:
data: Tuple that expands to:
a: torch.Tensor[float4e2m1fn] of shape [m, k, l],
b: torch.Tensor[float4e2m1fn] of shape [1, k, l],
sfa: torch.Tensor[float8_e4m3fnuz] of shape [m, k // 16, l], used by reference implementation
sfb: torch.Tensor[float8_e4m3fnuz] of shape [1, k // 16, l], used by reference implementation
sfa_permuted: torch.Tensor[float8_e4m3fnuz] of shape [32, 4, rest_m, 4, rest_k, l],
sfb_permuted: torch.Tensor[float8_e4m3fnuz] of shape [32, 4, rest_n, 4, rest_k, l],
c: torch.Tensor[float16] of shape [m, 1, l]
Returns:
Tensor containing output in float16
c: torch.Tensor[float16] of shape [m, 1, l]
"""
"""
PyTorch reference implementation of NVFP4 block-scaled GEMV.
"""
a_ref, b_ref, sfa, sfb, sfa_permuted, sfb_permuted, c_ref = data
# m, k, l = a_ref.shape
# if l == 1:
# tmp_c = torch.empty((l, 64, m), dtype=torch.float16, device="cuda")
# torch._scaled_mm(
# a_ref[:, :, 0],
# b_ref[0:64, :, 0].transpose(0, 1),
# sfa_permuted[:,:,:,:,:,0].permute(2, 4, 0, 1, 3).flatten(),
# sfb_permuted[:,:,:,:,:,0].permute(2, 4, 0, 1, 3).flatten(),
# bias=None,
# out_dtype=torch.float16,
# out=tmp_c[0,:,:].transpose(0,1),
# )
# return tmp_c[:, 0:1, :].permute(2, 1, 0)
# else:
# args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa_permuted.permute(5, 2, 4, 0, 1, 3), sfb_permuted.permute(5, 2, 4, 0, 1, 3), c_ref.permute(2, 0, 1).flatten())
args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa.permute(2, 0, 1), sfb.permute(2, 0, 1), c_ref.permute(2, 0, 1).flatten())
# for a in args:
# assert a.is_cuda, "All input tensors must be on GPU"
# print(a.stride(), a.shape)
module.custom_func(*args)
return c_refscrolls · 330 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 80103.
⋯ 16 unchanged linesusing fp8e4m3 = __nv_fp8_e4m3;using fp8e5m2 = __nv_fp8_e5m2;+ using fp8e4m3x4 = uint32_t;+using fp4e2m1 = __nv_fp4_e2m1;using fp4e2m1x2 = __nv_fp4x2_e2m1;using fp4e2m1x4 = __nv_fp4x4_e2m1;using e8m0_t = uint8_t;constexpr int WARP_SIZE = 32;- constexpr int WARP_COUNT = 16*2;+ constexpr int WARP_COUNT = 8;constexpr int THREADS_PER_BLOCK = WARP_SIZE * WARP_COUNT;constexpr int SM_COUNT = 148;⋯ 9 unchanged linesconstexpr int TILE_N = 8;constexpr int N_PADDED_128 = 128;constexpr int N = 1;+ constexpr int BYTES_PER_THREAD_PER_K_TILE = 32;-template <typename T>constexpr T DIVUP(const T &x, const T &y) {return (((x) + ((y)-1)) / (y));⋯ 18 unchanged linesuint64_t vec;fp4e2m1x2 small_vec[8];};-+ /* __device__ __forceinline__ fp8e4m3 pick_byte(uint32_t x, int byte_idx) {+ uint32_t tmp;+ // byte_idx in [0, 3], bit offset = byte_idx * 8, width = 8+ asm volatile (+ "bfe.u32 %0, %1, %2, 8;\n"+ : "=r"(tmp)+ : "r"(x), "r"(byte_idx * 8)+ );+ return static_cast<uint8_t>(tmp);+ } */__inline__ __device__uint32_t dot_fp4x8(const uint32_t a, const uint32_t b, uint32_t acc) {⋯ 33 unchanged lines}return val; // After this, lane 0 holds the sum of the warp}- __global__ void custom_func_kernel(const uint64_t* const __restrict__ A,- const uint64_t* const __restrict__ B,- const fp8e4m3 * const __restrict__ scale_A,+ __global__ void __maxnreg__(32) custom_func_kernel(const uint8_t* const __restrict__ A,+ const uint8_t* const __restrict__ B,+ fp8e4m3 * __restrict__ scale_A,const uint32_t * const __restrict__ scale_B,half* const __restrict__ C,const int L, const int M, const int K, const int sB_size_in_bytes) {⋯ 6 unchanged linesconst int lane_id = threadIdx.x % WARP_SIZE;int m_idx = block_m_idx * BLOCK_M + warp_id;- extern __shared__ __align__(16) uint64_t sB_u64[];- fp8e4m3 * sB_scale = reinterpret_cast<fp8e4m3 *>(- reinterpret_cast<uint8_t *>(sB_u64) + sB_size_in_bytes);+ extern __shared__ __align__(BYTES_PER_THREAD_PER_K_TILE) uint32_t sB_u32[];+ auto sB_scale = reinterpret_cast<fp8e4m3x4 *>(+ reinterpret_cast<uint8_t *>(sB_u32) + sB_size_in_bytes);// Load B into shared mem- int b_l_base_off = (l_idx * N_PADDED_128 * K) / (2 * sizeof(uint64_t));- int b_bound = K / (2 * sizeof(uint64_t));+ int b_l_base_off = (l_idx * N_PADDED_128 * K) / (2);+ int b_bound = K / (2 * BYTES_PER_THREAD_PER_K_TILE);- #pragma unrollfor (int i = threadIdx.x; i < b_bound; i += THREADS_PER_BLOCK) {- sB_u64[i] = B[b_l_base_off + i];+ uint32_t rB_values[8];+ {+ void *ptr = (void *)(&B[b_l_base_off + i*BYTES_PER_THREAD_PER_K_TILE]);+ asm volatile( \\+ "{\\n" \\+ "ld.global.v8.b32 { %0, %1, %2, %3, %4, %5, %6, %7 }, [%8];\\n" \\+ "}\\n" : "=r"(rB_values[0]), "=r"(rB_values[1]), "=r"(rB_values[2]), "=r"(rB_values[3]), "=r"(rB_values[4]), "=r"(rB_values[5]), "=r"(rB_values[6]), "=r"(rB_values[7]) : "l"(ptr));+ }+ int rest_idx = i / WARP_SIZE;++ #pragma unroll(8)+ for (int j = 0; j < 8; j++) {+ sB_u32[rest_idx * 256 + j * 32 + lane_id] = rB_values[j];+ }}int b_scale_bound = K / (16 * sizeof(uint32_t));+ auto curr_scale_B = &scale_B[(l_idx * N_PADDED_128 * K / (16 * sizeof(uint32_t)))];#pragma unrollfor (int i = threadIdx.x; i < b_scale_bound; i += THREADS_PER_BLOCK) {- reinterpret_cast<uint32_t *>(sB_scale)[i] = scale_B[scale_idx(0, i * sizeof(uint32_t), l_idx, N_PADDED_128, K) / sizeof(uint32_t)];+ reinterpret_cast<uint32_t *>(sB_scale)[i] = curr_scale_B[i];}⋯ 1 unchanged linesfor (;m_idx < M; m_idx += gridDim.x * BLOCK_M) {- const int a_l_base_off_u64 = (l_idx * (M * K) + m_idx * (K)) / (2 * sizeof(uint64_t));+ const int a_l_base_off = (l_idx * (M * K) + m_idx * (K)) / 2;constexpr int ATOM_SIZE = 128 * 4;const int TOTAL_NUM_ELEMENTS_PER_L = M * K / 16;const int rest_m = m_idx / 128;⋯ 1 unchanged linesconst int sub_m_idx0 = m_idx % 32;- const int scale_A_base = l_idx * (TOTAL_NUM_ELEMENTS_PER_L) + rest_m * (128 * K / 16) + sub_m_idx0 * 16 + sub_m_idx1 * 4;- const fp8e4m3 * const curr_scale_A = &scale_A[scale_A_base];- float acc{};- #pragma unroll- for (int tile_idx_k = lane_id; tile_idx_k < K / (2 * sizeof(uint64_t)); tile_idx_k += WARP_SIZE) {+ const int scale_A_base = l_idx * (TOTAL_NUM_ELEMENTS_PER_L) + m_idx * (K / 16);+ auto curr_scale_A = reinterpret_cast<fp8e4m3x4 * >(&scale_A[scale_A_base]);+ half acc{};+ for (int tile_idx_k = lane_id; tile_idx_k < K / (2 * BYTES_PER_THREAD_PER_K_TILE); tile_idx_k += WARP_SIZE) {+ constexpr int factor = (2 * BYTES_PER_THREAD_PER_K_TILE) / (16 * sizeof(fp8e4m3x4));+ uint32_t a_scales = curr_scale_A[tile_idx_k * factor];+ uint32_t b_scales = sB_scale[tile_idx_k];+ union {+ uint32_t u;+ fp8e4m3 f[4];+ } cvt_a, cvt_b;+ cvt_a.u = a_scales;+ cvt_b.u = b_scales;+ uint32_t rA_values[8];- uint64_t rA_value = A[a_l_base_off_u64 + tile_idx_k];- uint64_t rB_value = sB_u64[tile_idx_k];+ {+ void *ptr = (void *)(&A[a_l_base_off + tile_idx_k*BYTES_PER_THREAD_PER_K_TILE]);+ asm volatile( \\+ "{\\n" \\+ "ld.global.v8.b32 { %0, %1, %2, %3, %4, %5, %6, %7 }, [%8];\\n" \\+ "}\\n" : "=r"(rA_values[0]), "=r"(rA_values[1]), "=r"(rA_values[2]), "=r"(rA_values[3]), "=r"(rA_values[4]), "=r"(rA_values[5]), "=r"(rA_values[6]), "=r"(rA_values[7]) : "l"(ptr));+ }- uint32_t tile_acc2{};-+ const int base_off = (tile_idx_k / WARP_SIZE)*256 + lane_id;-- tile_acc2 = dot_fp4x8(static_cast<uint32_t>(rA_value&0xffffffff), static_cast<uint32_t>(rB_value&0xffffffff), tile_acc2);- tile_acc2 = dot_fp4x8(static_cast<uint32_t>(rA_value>>32), static_cast<uint32_t>(rB_value>>32), tile_acc2);- half2 tile_acc = reinterpret_cast<__half2 &>(tile_acc2);+ #pragma unroll(4)+ for (int i = 0; i < 8; i+=2) {+ uint32_t tile_acc2{};+ tile_acc2 = dot_fp4x8(rA_values[i], sB_u32[base_off + 32 * i], tile_acc2);+ tile_acc2 = dot_fp4x8(rA_values[i+1], sB_u32[base_off + 32 * (i + 1)], tile_acc2);+ half2 tile_acc = reinterpret_cast<__half2 &>(tile_acc2);+ acc += static_cast<half>(tile_acc.x + tile_acc.y) * static_cast<half>(cvt_a.f[i>>1]) * static_cast<half>(cvt_b.f[i>>1]);- const int rest_k = tile_idx_k >> 2;- const int sub_k_idx = tile_idx_k & 3;- const float a_scale = static_cast<float>(curr_scale_A[rest_k * ATOM_SIZE + sub_k_idx]);--- const float b_scale = static_cast<float>(sB_scale[tile_idx_k]);+ }- acc += static_cast<float>(tile_acc.x + tile_acc.y) * a_scale * b_scale;-}acc = warpReduceSum(acc);if (lane_id == 0) {⋯ 11 unchanged linesconst int threads = THREADS_PER_BLOCK;// const int blocks = (N + threads - 1) / threads;-- const int sB_size_in_bytes = K / 2;+ constexpr int sB_block_size = 256*4;+ const int sB_size_in_bytes = cuda::ceil_div(K / 2, sB_block_size) * sB_block_size;const int sB_scale_size_in_bytes = K / 16;const int shared_mem_size_in_bytes = sB_size_in_bytes + sB_scale_size_in_bytes;⋯ 1 unchanged linesconst int GRID_DIM_X = cuda::ceil_div(M, BLOCK_M);dim3 blocks(GRID_DIM_X, L, 1);custom_func_kernel<<<blocks, threads, shared_mem_size_in_bytes>>>(- reinterpret_cast<uint64_t *>(A.data_ptr()),- reinterpret_cast<uint64_t *>(B.data_ptr()),+ reinterpret_cast<uint8_t *>(A.data_ptr()),+ reinterpret_cast<uint8_t *>(B.data_ptr()),reinterpret_cast<fp8e4m3 *>(scale_A.data_ptr()),reinterpret_cast<uint32_t *>(scale_B.data_ptr()),reinterpret_cast<half *>(C.data_ptr()),⋯ 53 unchanged lines"""PyTorch reference implementation of NVFP4 block-scaled GEMV."""- a_ref, b_ref, _, _, sfa_permuted, sfb_permuted, c_ref = data+ a_ref, b_ref, sfa, sfb, sfa_permuted, sfb_permuted, c_ref = data# m, k, l = a_ref.shape# if l == 1:⋯ 10 unchanged lines# )# return tmp_c[:, 0:1, :].permute(2, 1, 0)# else:- args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa_permuted.permute(5, 2, 4, 0, 1, 3), sfb_permuted.permute(5, 2, 4, 0, 1, 3), c_ref.permute(2, 0, 1).flatten())-+ # args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa_permuted.permute(5, 2, 4, 0, 1, 3), sfb_permuted.permute(5, 2, 4, 0, 1, 3), c_ref.permute(2, 0, 1).flatten())+ args = (a_ref.permute(2, 0, 1), b_ref.permute(2, 0, 1), sfa.permute(2, 0, 1), sfb.permute(2, 0, 1), c_ref.permute(2, 0, 1).flatten())+ # for a in args:+ # assert a.is_cuda, "All input tensors must be on GPU"+ # print(a.stride(), a.shape)module.custom_func(*args)return c_refNo newline at end of file
scrolls · 226 diff lines total
Best evidence level for this revision: reported
JSON