submission 113565
macto · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 491 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-113565?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:cc0c0f48b6108ba69fcfc4ee1a660d06b9cc6045b581d2438c3e91f601cbe422
license declaredunknown
license concludedunknown
authorsmacto
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" \shared-memory
__shared__ __align__(16) uint8_t sh_a[2][ROWS_PER_BLOCK][K_TILE_BYTES];vector-width = half2
__device__ __forceinline__ half2 cvt_fp8_to_h2(uint16_t x) {Kernel source
submission.py491 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
CUDA_SRC = r'''
#include <torch/extension.h>
#include <cuda_fp16.h>
#include <cuda_fp4.h>
#include <cuda_fp8.h>
enum CachePolicy { CS_CA, LU_CA, NC_EVICT };
template<CachePolicy P>
__device__ __forceinline__ void load_vec4(uint32_t* dst, const void* src) {
if constexpr (P == NC_EVICT) {
asm volatile("ld.global.nc.L1::no_allocate.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
} else if constexpr (P == LU_CA) {
asm volatile("ld.global.lu.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
} else {
asm volatile("ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
}
}
template<CachePolicy P>
__device__ __forceinline__ void load_vec4_b(uint32_t* dst, const void* src) {
if constexpr (P == NC_EVICT) {
asm volatile("ld.global.nc.L1::evict_last.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
} else {
asm volatile("ld.global.ca.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
}
}
__device__ __forceinline__ void load_vec4_b_l2(uint32_t* dst, const void* src) {
asm volatile("ld.global.L2::128B.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));
}
template<CachePolicy P>
__device__ __forceinline__ void load_u16(uint16_t* dst, const void* src) {
if constexpr (P == NC_EVICT) {
asm volatile("ld.global.nc.L1::no_allocate.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
} else if constexpr (P == LU_CA) {
asm volatile("ld.global.lu.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
} else {
asm volatile("ld.global.cs.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
}
}
template<CachePolicy P>
__device__ __forceinline__ void load_u16_b(uint16_t* dst, const void* src) {
if constexpr (P == NC_EVICT) {
asm volatile("ld.global.nc.L1::evict_last.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
} else {
asm volatile("ld.global.ca.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
}
}
__device__ __forceinline__ void load_u16_b_l2(uint16_t* dst, const void* src) {
asm volatile("ld.global.L2::128B.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
}
__device__ __forceinline__ void load_u32_lu(uint32_t* dst, const void* src) {
asm volatile("ld.global.lu.u32 %0, [%1];" : "=r"(*dst) : "l"(src));
}
__device__ __forceinline__ void load_u32_b_l2(uint32_t* dst, const void* src) {
asm volatile("ld.global.L2::128B.u32 %0, [%1];" : "=r"(*dst) : "l"(src));
}
__device__ __forceinline__ half2 cvt_fp8_to_h2(uint16_t x) {
int out;
asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(x));
return *reinterpret_cast<half2*>(&out);
}
__device__ __forceinline__ void cvt_fp8x4_to_h2x4(half2* out, uint32_t in) {
asm volatile(
"{\n\t"
".reg .b16 lo, hi;\n\t"
"mov.b32 {lo, hi}, %2;\n\t"
"cvt.rn.f16x2.e4m3x2 %0, lo;\n\t"
"cvt.rn.f16x2.e4m3x2 %1, hi;\n\t"
"}"
: "=r"(reinterpret_cast<int&>(out[0])),
"=r"(reinterpret_cast<int&>(out[1]))
: "r"(in)
);
}
__device__ __forceinline__ void cvt_fp4x4_to_h2x4(half2* out, uint32_t in) {
asm volatile(
"{\n\t"
".reg .b8 b0, b1, b2, b3;\n\t"
"mov.b32 {b0, b1, b2, b3}, %4;\n\t"
"cvt.rn.f16x2.e2m1x2 %0, b0;\n\t"
"cvt.rn.f16x2.e2m1x2 %1, b1;\n\t"
"cvt.rn.f16x2.e2m1x2 %2, b2;\n\t"
"cvt.rn.f16x2.e2m1x2 %3, b3;\n\t"
"}"
: "=r"(reinterpret_cast<int&>(out[0])),
"=r"(reinterpret_cast<int&>(out[1])),
"=r"(reinterpret_cast<int&>(out[2])),
"=r"(reinterpret_cast<int&>(out[3]))
: "r"(in)
);
}
__device__ __forceinline__ float dot_scaled_32fp4(
const uint32_t* a, const uint32_t* b, uint16_t sfa, uint16_t sfb
) {
half2 scale = __hmul2(cvt_fp8_to_h2(sfa), cvt_fp8_to_h2(sfb));
half2 a_h[16], b_h[16];
#pragma unroll
for (int i = 0; i < 4; i++) {
cvt_fp4x4_to_h2x4(a_h + i * 4, a[i]);
cvt_fp4x4_to_h2x4(b_h + i * 4, b[i]);
}
half2 acc0 = __hmul2(a_h[0], b_h[0]);
half2 acc1 = __hmul2(a_h[8], b_h[8]);
#pragma unroll
for (int i = 1; i < 8; i++) {
acc0 = __hfma2(a_h[i], b_h[i], acc0);
acc1 = __hfma2(a_h[8 + i], b_h[8 + i], acc1);
}
half sum0 = __hadd(acc0.x, acc0.y);
half sum1 = __hadd(acc1.x, acc1.y);
return __half2float(__hadd(__hmul(sum0, scale.x), __hmul(sum1, scale.y)));
}
__device__ __forceinline__ float dot_scaled_64fp4(
const uint32_t* a, const uint32_t* b, uint32_t sfa, uint32_t sfb
) {
half2 sfa_h[2], sfb_h[2];
cvt_fp8x4_to_h2x4(sfa_h, sfa);
cvt_fp8x4_to_h2x4(sfb_h, sfb);
half2 scale0 = __hmul2(sfa_h[0], sfb_h[0]);
half2 scale1 = __hmul2(sfa_h[1], sfb_h[1]);
half2 a_h[32], b_h[32];
#pragma unroll
for (int i = 0; i < 8; i++) {
cvt_fp4x4_to_h2x4(a_h + i * 4, a[i]);
cvt_fp4x4_to_h2x4(b_h + i * 4, b[i]);
}
half2 acc0 = __hmul2(a_h[0], b_h[0]);
half2 acc1 = __hmul2(a_h[8], b_h[8]);
half2 acc2 = __hmul2(a_h[16], b_h[16]);
half2 acc3 = __hmul2(a_h[24], b_h[24]);
#pragma unroll
for (int i = 1; i < 8; i++) {
acc0 = __hfma2(a_h[i], b_h[i], acc0);
acc1 = __hfma2(a_h[8 + i], b_h[8 + i], acc1);
acc2 = __hfma2(a_h[16 + i], b_h[16 + i], acc2);
acc3 = __hfma2(a_h[24 + i], b_h[24 + i], acc3);
}
half sum0 = __hadd(acc0.x, acc0.y);
half sum1 = __hadd(acc1.x, acc1.y);
half sum2 = __hadd(acc2.x, acc2.y);
half sum3 = __hadd(acc3.x, acc3.y);
half scaled0 = __hadd(__hmul(sum0, scale0.x), __hmul(sum1, scale0.y));
half scaled1 = __hadd(__hmul(sum2, scale1.x), __hmul(sum3, scale1.y));
return __half2float(__hadd(scaled0, scaled1));
}
template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CachePolicy POLICY>
__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)
void gemv_direct_kernel(
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, int N_pad
) {
const int row = blockIdx.x * ROWS_PER_BLOCK + threadIdx.x / THREADS_PER_ROW;
const int batch = blockIdx.y;
const int lane = threadIdx.x % THREADS_PER_ROW;
if (row >= M) return;
const int K_bytes = K / 2;
const int K_sf = K / 16;
const uint8_t* a_ptr = a + (size_t)batch * M * K_bytes + row * K_bytes;
const uint8_t* b_ptr = b + (size_t)batch * N_pad * K_bytes;
const uint8_t* sfa_ptr = sfa + (size_t)batch * M * K_sf + row * K_sf;
const uint8_t* sfb_ptr = sfb + (size_t)batch * N_pad * K_sf;
float acc = 0.0f;
const int num_sf_pairs = K_sf / 2;
#pragma unroll 4
for (int sp = lane; sp < num_sf_pairs; sp += THREADS_PER_ROW) {
const int sf_idx = sp * 2;
const int byte_idx = sf_idx * 8;
uint16_t sf_a, sf_b;
load_u16<POLICY>(&sf_a, sfa_ptr + sf_idx);
load_u16_b<POLICY>(&sf_b, sfb_ptr + sf_idx);
uint32_t a_regs[4], b_regs[4];
load_vec4<POLICY>(a_regs, a_ptr + byte_idx);
load_vec4_b<POLICY>(b_regs, b_ptr + byte_idx);
acc += dot_scaled_32fp4(a_regs, b_regs, sf_a, sf_b);
}
#pragma unroll
for (int off = THREADS_PER_ROW / 2; off > 0; off /= 2)
acc += __shfl_down_sync(0xffffffff, acc, off);
if (lane == 0)
c[(size_t)row + (size_t)batch * M] = __float2half(acc);
}
template<int ROWS_PER_BLOCK, int THREADS_PER_ROW>
__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)
void gemv_wide_l2_kernel(
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, int N_pad
) {
const int row = blockIdx.x * ROWS_PER_BLOCK + threadIdx.x / THREADS_PER_ROW;
const int batch = blockIdx.y;
const int lane = threadIdx.x % THREADS_PER_ROW;
if (row >= M) return;
const int K_bytes = K / 2;
const int K_sf = K / 16;
const uint8_t* a_ptr = a + (size_t)batch * M * K_bytes + row * K_bytes;
const uint8_t* b_ptr = b + (size_t)batch * N_pad * K_bytes;
const uint8_t* sfa_ptr = sfa + (size_t)batch * M * K_sf + row * K_sf;
const uint8_t* sfb_ptr = sfb + (size_t)batch * N_pad * K_sf;
float acc = 0.0f;
const int num_sf_quads = K_sf / 4;
#pragma unroll 2
for (int sq = lane; sq < num_sf_quads; sq += THREADS_PER_ROW) {
const int sf_idx = sq * 4;
const int byte_idx = sf_idx * 8;
uint32_t sf_a, sf_b;
load_u32_lu(&sf_a, sfa_ptr + sf_idx);
load_u32_b_l2(&sf_b, sfb_ptr + sf_idx);
uint32_t a_regs[8], b_regs[8];
asm volatile("ld.global.lu.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(a_regs[0]), "=r"(a_regs[1]), "=r"(a_regs[2]), "=r"(a_regs[3])
: "l"(a_ptr + byte_idx));
asm volatile("ld.global.lu.v4.u32 {%0,%1,%2,%3}, [%4];"
: "=r"(a_regs[4]), "=r"(a_regs[5]), "=r"(a_regs[6]), "=r"(a_regs[7])
: "l"(a_ptr + byte_idx + 16));
load_vec4_b_l2(b_regs, b_ptr + byte_idx);
load_vec4_b_l2(b_regs + 4, b_ptr + byte_idx + 16);
acc += dot_scaled_64fp4(a_regs, b_regs, sf_a, sf_b);
}
#pragma unroll
for (int off = THREADS_PER_ROW / 2; off > 0; off /= 2)
acc += __shfl_down_sync(0xffffffff, acc, off);
if (lane == 0)
c[(size_t)row + (size_t)batch * M] = __float2half(acc);
}
#define ASYNC_CP_16(dst, src) \
asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" \
:: "r"((uint32_t)__cvta_generic_to_shared(dst)), "l"(src))
#define ASYNC_CP_4(dst, src) \
asm volatile("cp.async.ca.shared.global [%0], [%1], 4;" \
:: "r"((uint32_t)__cvta_generic_to_shared(dst)), "l"(src))
#define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")
#define ASYNC_WAIT(n) asm volatile("cp.async.wait_group %0;" :: "n"(n))
template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, int K_TILE>
__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)
void gemv_pipelined_kernel(
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, int N_pad
) {
constexpr int K_TILE_BYTES = K_TILE / 2;
constexpr int K_TILE_SF = K_TILE / 16;
__shared__ __align__(16) uint8_t sh_a[2][ROWS_PER_BLOCK][K_TILE_BYTES];
__shared__ __align__(16) uint8_t sh_b[2][K_TILE_BYTES];
__shared__ __align__(16) uint8_t sh_sfa[2][ROWS_PER_BLOCK][K_TILE_SF];
__shared__ __align__(16) uint8_t sh_sfb[2][K_TILE_SF];
const int row_base = blockIdx.x * ROWS_PER_BLOCK;
const int batch = blockIdx.y;
const int tid = threadIdx.x;
const int warp_id = tid / THREADS_PER_ROW;
const int lane = tid % THREADS_PER_ROW;
const int row = row_base + warp_id;
if (row >= M) return;
const int K_bytes = K / 2;
const int K_sf = K / 16;
const int num_tiles = (K_bytes + K_TILE_BYTES - 1) / K_TILE_BYTES;
const size_t a_base = (size_t)batch * M * K_bytes;
const size_t b_base = (size_t)batch * N_pad * K_bytes;
const size_t sfa_base = (size_t)batch * M * K_sf;
const size_t sfb_base = (size_t)batch * N_pad * K_sf;
auto issue_tile = [&](int tile, int stage) {
const int k_off_bytes = tile * K_TILE_BYTES;
const int k_off_sf = tile * K_TILE_SF;
const int bytes_this = min(K_TILE_BYTES, K_bytes - k_off_bytes);
const int sf_this = min(K_TILE_SF, K_sf - k_off_sf);
if (warp_id < ROWS_PER_BLOCK && (row_base + warp_id) < M) {
const uint8_t* a_row = a + a_base + (row_base + warp_id) * K_bytes + k_off_bytes;
for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16)
if (i + 16 <= bytes_this) ASYNC_CP_16(&sh_a[stage][warp_id][i], a_row + i);
const uint8_t* sfa_row = sfa + sfa_base + (row_base + warp_id) * K_sf + k_off_sf;
for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4)
if (i + 4 <= sf_this) ASYNC_CP_4(&sh_sfa[stage][warp_id][i], sfa_row + i);
}
if (warp_id == 0) {
const uint8_t* b_ptr = b + b_base + k_off_bytes;
for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16)
if (i + 16 <= bytes_this) ASYNC_CP_16(&sh_b[stage][i], b_ptr + i);
const uint8_t* sfb_ptr = sfb + sfb_base + k_off_sf;
for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4)
if (i + 4 <= sf_this) ASYNC_CP_4(&sh_sfb[stage][i], sfb_ptr + i);
}
ASYNC_COMMIT();
};
issue_tile(0, 0);
ASYNC_WAIT(0);
__syncthreads();
float acc = 0.0f;
for (int tile = 0; tile < num_tiles; tile++) {
const int stage = tile & 1;
if (tile + 1 < num_tiles)
issue_tile(tile + 1, (tile + 1) & 1);
const int k_off_sf = tile * K_TILE_SF;
const int sf_this = min(K_TILE_SF, K_sf - k_off_sf);
const int num_sf_pairs = sf_this / 2;
#pragma unroll 2
for (int sp = lane; sp < num_sf_pairs; sp += THREADS_PER_ROW) {
const int sf_idx = sp * 2;
const int byte_idx = sf_idx * 8;
uint16_t sf_a = *reinterpret_cast<const uint16_t*>(&sh_sfa[stage][warp_id][sf_idx]);
uint16_t sf_b = *reinterpret_cast<const uint16_t*>(&sh_sfb[stage][sf_idx]);
const uint32_t* ap = reinterpret_cast<const uint32_t*>(&sh_a[stage][warp_id][byte_idx]);
const uint32_t* bp = reinterpret_cast<const uint32_t*>(&sh_b[stage][byte_idx]);
acc += dot_scaled_32fp4(ap, bp, sf_a, sf_b);
}
if (tile + 1 < num_tiles) {
ASYNC_WAIT(0);
__syncthreads();
}
}
#pragma unroll
for (int off = THREADS_PER_ROW / 2; off > 0; off /= 2)
acc += __shfl_down_sync(0xffffffff, acc, off);
if (lane == 0)
c[(size_t)row + (size_t)batch * M] = __float2half(acc);
}
void run_gemv(
torch::Tensor C, torch::Tensor A, torch::Tensor B,
torch::Tensor SFA, torch::Tensor SFB,
int m, int k, int l, int n_pad
) {
auto a_ptr = A.data_ptr<uint8_t>();
auto b_ptr = B.data_ptr<uint8_t>();
auto sfa_ptr = SFA.data_ptr<uint8_t>();
auto sfb_ptr = SFB.data_ptr<uint8_t>();
auto c_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());
bool use_pipelined = (k >= 8192 && l <= 2) || (k <= 4096 && l >= 4);
if (use_pipelined) {
dim3 grid((m + 3) / 4, l);
dim3 block(128);
if (k >= 8192) {
gemv_pipelined_kernel<4, 32, 4096><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
} else {
gemv_pipelined_kernel<4, 32, 2048><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
}
}
else if (l >= 8) {
dim3 grid((m + 7) / 8, l);
dim3 block(128);
gemv_wide_l2_kernel<8, 16><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
}
else if (l >= 4 && k >= 4096) {
dim3 grid((m + 7) / 8, l);
dim3 block(128);
gemv_direct_kernel<8, 16, LU_CA><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
}
else {
dim3 grid((m + 3) / 4, l);
dim3 block(128);
gemv_direct_kernel<4, 32, CS_CA><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
}
}
'''
CPP_SRC = r'''
#include <torch/extension.h>
void run_gemv(torch::Tensor C, torch::Tensor A, torch::Tensor B,
torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad);
'''
_module = None
def _load_module():
global _module
if _module is None:
_module = load_inline(
name='nvfp4_gemv_v24',
cpp_sources=CPP_SRC,
cuda_sources=CUDA_SRC,
functions=['run_gemv'],
extra_cuda_cflags=[
'-O3', '--use_fast_math', '-std=c++17',
'-gencode=arch=compute_100a,code=sm_100a',
],
verbose=False,
)
return _module
def custom_kernel(data: input_t) -> output_t:
a, b, sfa, sfb, _, _, c = data
m, k_packed, l = a.shape
k = k_packed * 2
n_pad = b.shape[0]
if not sfa.is_cuda:
sfa = sfa.to(a.device)
if not sfb.is_cuda:
sfb = sfb.to(a.device)
c_out = c.squeeze(1)
_load_module().run_gemv(
c_out,
a.view(torch.uint8),
b.view(torch.uint8),
sfa.view(torch.uint8),
sfb.view(torch.uint8),
m, k, l, n_pad
)
return c
scrolls · 491 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 110022.
- """- Dual-module hybrid: Loads direct and pipelined kernels as SEPARATE modules- to avoid any compiler optimization interference between them.- """import torchfrom torch.utils.cpp_extension import load_inlinefrom task import input_t, output_t- # =============================================================================- # DIRECT KERNEL MODULE (exact copy from v13)- # =============================================================================-- DIRECT_CUDA = r'''+ CUDA_SRC = r'''#include <torch/extension.h>#include <cuda_fp16.h>#include <cuda_fp4.h>#include <cuda_fp8.h>- enum class CacheHint { CS, CA };+ enum CachePolicy { CS_CA, LU_CA, NC_EVICT };- template<CacheHint hint>+ template<CachePolicy P>__device__ __forceinline__ void load_vec4(uint32_t* dst, const void* src) {- if constexpr (hint == CacheHint::CS) {- asm volatile("ld.global.cs.v4.u32 {%0, %1, %2, %3}, [%4];"+ if constexpr (P == NC_EVICT) {+ asm volatile("ld.global.nc.L1::no_allocate.v4.u32 {%0,%1,%2,%3}, [%4];": "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));+ } else if constexpr (P == LU_CA) {+ asm volatile("ld.global.lu.v4.u32 {%0,%1,%2,%3}, [%4];"+ : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));} else {- asm volatile("ld.global.ca.v4.u32 {%0, %1, %2, %3}, [%4];"+ asm volatile("ld.global.cs.v4.u32 {%0,%1,%2,%3}, [%4];": "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));}}- template<CacheHint hint>+ template<CachePolicy P>+ __device__ __forceinline__ void load_vec4_b(uint32_t* dst, const void* src) {+ if constexpr (P == NC_EVICT) {+ asm volatile("ld.global.nc.L1::evict_last.v4.u32 {%0,%1,%2,%3}, [%4];"+ : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));+ } else {+ asm volatile("ld.global.ca.v4.u32 {%0,%1,%2,%3}, [%4];"+ : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));+ }+ }++ __device__ __forceinline__ void load_vec4_b_l2(uint32_t* dst, const void* src) {+ asm volatile("ld.global.L2::128B.v4.u32 {%0,%1,%2,%3}, [%4];"+ : "=r"(dst[0]), "=r"(dst[1]), "=r"(dst[2]), "=r"(dst[3]) : "l"(src));+ }++ template<CachePolicy P>__device__ __forceinline__ void load_u16(uint16_t* dst, const void* src) {- if constexpr (hint == CacheHint::CS) {+ if constexpr (P == NC_EVICT) {+ asm volatile("ld.global.nc.L1::no_allocate.u16 %0, [%1];" : "=h"(*dst) : "l"(src));+ } else if constexpr (P == LU_CA) {+ asm volatile("ld.global.lu.u16 %0, [%1];" : "=h"(*dst) : "l"(src));+ } else {asm volatile("ld.global.cs.u16 %0, [%1];" : "=h"(*dst) : "l"(src));+ }+ }++ template<CachePolicy P>+ __device__ __forceinline__ void load_u16_b(uint16_t* dst, const void* src) {+ if constexpr (P == NC_EVICT) {+ asm volatile("ld.global.nc.L1::evict_last.u16 %0, [%1];" : "=h"(*dst) : "l"(src));} else {asm volatile("ld.global.ca.u16 %0, [%1];" : "=h"(*dst) : "l"(src));}}- __device__ __forceinline__ void cvt_fp4_to_f16(half2* out, uint32_t in) {+ __device__ __forceinline__ void load_u16_b_l2(uint16_t* dst, const void* src) {+ asm volatile("ld.global.L2::128B.u16 %0, [%1];" : "=h"(*dst) : "l"(src));+ }++ __device__ __forceinline__ void load_u32_lu(uint32_t* dst, const void* src) {+ asm volatile("ld.global.lu.u32 %0, [%1];" : "=r"(*dst) : "l"(src));+ }++ __device__ __forceinline__ void load_u32_b_l2(uint32_t* dst, const void* src) {+ asm volatile("ld.global.L2::128B.u32 %0, [%1];" : "=r"(*dst) : "l"(src));+ }++ __device__ __forceinline__ half2 cvt_fp8_to_h2(uint16_t x) {+ int out;+ asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(x));+ return *reinterpret_cast<half2*>(&out);+ }++ __device__ __forceinline__ void cvt_fp8x4_to_h2x4(half2* out, uint32_t in) {asm volatile("{\n\t"+ ".reg .b16 lo, hi;\n\t"+ "mov.b32 {lo, hi}, %2;\n\t"+ "cvt.rn.f16x2.e4m3x2 %0, lo;\n\t"+ "cvt.rn.f16x2.e4m3x2 %1, hi;\n\t"+ "}"+ : "=r"(reinterpret_cast<int&>(out[0])),+ "=r"(reinterpret_cast<int&>(out[1]))+ : "r"(in)+ );+ }++ __device__ __forceinline__ void cvt_fp4x4_to_h2x4(half2* out, uint32_t in) {+ asm volatile(+ "{\n\t"".reg .b8 b0, b1, b2, b3;\n\t""mov.b32 {b0, b1, b2, b3}, %4;\n\t""cvt.rn.f16x2.e2m1x2 %0, b0;\n\t"⋯ 9 unchanged lines);}- __device__ __forceinline__ half2 cvt_fp8_to_f16(uint16_t in) {- int out;- asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(in));- return *reinterpret_cast<half2*>(&out);- }-__device__ __forceinline__ float dot_scaled_32fp4(- const uint32_t* a_data,- const uint32_t* b_data,- half2 scale+ const uint32_t* a, const uint32_t* b, uint16_t sfa, uint16_t sfb) {- half2 a_f16[16], b_f16[16];+ half2 scale = __hmul2(cvt_fp8_to_h2(sfa), cvt_fp8_to_h2(sfb));++ half2 a_h[16], b_h[16];+#pragma unrollfor (int i = 0; i < 4; i++) {- cvt_fp4_to_f16(a_f16 + i * 4, a_data[i]);- cvt_fp4_to_f16(b_f16 + i * 4, b_data[i]);+ cvt_fp4x4_to_h2x4(a_h + i * 4, a[i]);+ cvt_fp4x4_to_h2x4(b_h + i * 4, b[i]);}- half2 acc0 = __hmul2(a_f16[0], b_f16[0]);- half2 acc1 = __hmul2(a_f16[8], b_f16[8]);+ half2 acc0 = __hmul2(a_h[0], b_h[0]);+ half2 acc1 = __hmul2(a_h[8], b_h[8]);#pragma unrollfor (int i = 1; i < 8; i++) {- acc0 = __hfma2(a_f16[i], b_f16[i], acc0);- acc1 = __hfma2(a_f16[8 + i], b_f16[8 + i], acc1);+ acc0 = __hfma2(a_h[i], b_h[i], acc0);+ acc1 = __hfma2(a_h[8 + i], b_h[8 + i], acc1);}half sum0 = __hadd(acc0.x, acc0.y);half sum1 = __hadd(acc1.x, acc1.y);- float result = 0.0f;- asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"- : "+f"(result) : "h"(*(uint16_t*)&sum0), "h"(*(uint16_t*)&scale.x));- asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"- : "+f"(result) : "h"(*(uint16_t*)&sum1), "h"(*(uint16_t*)&scale.y));+ return __half2float(__hadd(__hmul(sum0, scale.x), __hmul(sum1, scale.y)));+ }++ __device__ __forceinline__ float dot_scaled_64fp4(+ const uint32_t* a, const uint32_t* b, uint32_t sfa, uint32_t sfb+ ) {+ half2 sfa_h[2], sfb_h[2];+ cvt_fp8x4_to_h2x4(sfa_h, sfa);+ cvt_fp8x4_to_h2x4(sfb_h, sfb);- return result;+ half2 scale0 = __hmul2(sfa_h[0], sfb_h[0]);+ half2 scale1 = __hmul2(sfa_h[1], sfb_h[1]);++ half2 a_h[32], b_h[32];++ #pragma unroll+ for (int i = 0; i < 8; i++) {+ cvt_fp4x4_to_h2x4(a_h + i * 4, a[i]);+ cvt_fp4x4_to_h2x4(b_h + i * 4, b[i]);+ }++ half2 acc0 = __hmul2(a_h[0], b_h[0]);+ half2 acc1 = __hmul2(a_h[8], b_h[8]);+ half2 acc2 = __hmul2(a_h[16], b_h[16]);+ half2 acc3 = __hmul2(a_h[24], b_h[24]);++ #pragma unroll+ for (int i = 1; i < 8; i++) {+ acc0 = __hfma2(a_h[i], b_h[i], acc0);+ acc1 = __hfma2(a_h[8 + i], b_h[8 + i], acc1);+ acc2 = __hfma2(a_h[16 + i], b_h[16 + i], acc2);+ acc3 = __hfma2(a_h[24 + i], b_h[24 + i], acc3);+ }++ half sum0 = __hadd(acc0.x, acc0.y);+ half sum1 = __hadd(acc1.x, acc1.y);+ half sum2 = __hadd(acc2.x, acc2.y);+ half sum3 = __hadd(acc3.x, acc3.y);++ half scaled0 = __hadd(__hmul(sum0, scale0.x), __hmul(sum1, scale0.y));+ half scaled1 = __hadd(__hmul(sum2, scale1.x), __hmul(sum3, scale1.y));++ return __half2float(__hadd(scaled0, scaled1));}- template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CacheHint A_HINT, CacheHint B_HINT>+ template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CachePolicy POLICY>__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)- void gemv_kernel(- 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, int N_pad+ void gemv_direct_kernel(+ 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, int N_pad) {- const int row_base = blockIdx.x * ROWS_PER_BLOCK;+ const int row = blockIdx.x * ROWS_PER_BLOCK + threadIdx.x / THREADS_PER_ROW;const int batch = blockIdx.y;- const int tid = threadIdx.x;+ const int lane = threadIdx.x % THREADS_PER_ROW;- const int warp_id = tid / THREADS_PER_ROW;- const int lane = tid % THREADS_PER_ROW;-- const int row = row_base + warp_id;if (row >= M) return;const int K_bytes = K / 2;const int K_sf = K / 16;- const size_t a_off = (size_t)batch * M * K_bytes + row * K_bytes;- const size_t b_off = (size_t)batch * N_pad * K_bytes;- const size_t sfa_off = (size_t)batch * M * K_sf + row * K_sf;- const size_t sfb_off = (size_t)batch * N_pad * K_sf;+ const uint8_t* a_ptr = a + (size_t)batch * M * K_bytes + row * K_bytes;+ const uint8_t* b_ptr = b + (size_t)batch * N_pad * K_bytes;+ const uint8_t* sfa_ptr = sfa + (size_t)batch * M * K_sf + row * K_sf;+ const uint8_t* sfb_ptr = sfb + (size_t)batch * N_pad * K_sf;- const uint8_t* a_ptr = a + a_off;- const uint8_t* b_ptr = b + b_off;- const uint8_t* sfa_ptr = sfa + sfa_off;- const uint8_t* sfb_ptr = sfb + sfb_off;-float acc = 0.0f;const int num_sf_pairs = K_sf / 2;⋯ 2 unchanged linesconst int sf_idx = sp * 2;const int byte_idx = sf_idx * 8;- uint16_t sfa_packed, sfb_packed;- load_u16<A_HINT>(&sfa_packed, sfa_ptr + sf_idx);- load_u16<B_HINT>(&sfb_packed, sfb_ptr + sf_idx);- half2 scale = __hmul2(cvt_fp8_to_f16(sfa_packed), cvt_fp8_to_f16(sfb_packed));+ uint16_t sf_a, sf_b;+ load_u16<POLICY>(&sf_a, sfa_ptr + sf_idx);+ load_u16_b<POLICY>(&sf_b, sfb_ptr + sf_idx);uint32_t a_regs[4], b_regs[4];- load_vec4<A_HINT>(a_regs, a_ptr + byte_idx);- load_vec4<B_HINT>(b_regs, b_ptr + byte_idx);+ load_vec4<POLICY>(a_regs, a_ptr + byte_idx);+ load_vec4_b<POLICY>(b_regs, b_ptr + byte_idx);- acc += dot_scaled_32fp4(a_regs, b_regs, scale);+ acc += dot_scaled_32fp4(a_regs, b_regs, sf_a, sf_b);}#pragma unroll- for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {- acc += __shfl_down_sync(0xffffffff, acc, offset);- }+ for (int off = THREADS_PER_ROW / 2; off > 0; off /= 2)+ acc += __shfl_down_sync(0xffffffff, acc, off);- if (lane == 0) {+ if (lane == 0)c[(size_t)row + (size_t)batch * M] = __float2half(acc);- }}- template<int ROWS, int THREADS, CacheHint A_HINT, CacheHint B_HINT>- void dispatch(torch::Tensor& c, torch::Tensor& a, torch::Tensor& b,- torch::Tensor& sfa, torch::Tensor& sfb, int m, int k, int l, int n) {- dim3 grid((m + ROWS - 1) / ROWS, l);- dim3 block(ROWS * THREADS);- gemv_kernel<ROWS, THREADS, A_HINT, B_HINT><<<grid, block>>>(- a.data_ptr<uint8_t>(), b.data_ptr<uint8_t>(),- sfa.data_ptr<uint8_t>(), sfb.data_ptr<uint8_t>(),- reinterpret_cast<half*>(c.data_ptr<at::Half>()), m, k, l, n);- }-- void run_direct(- torch::Tensor C, torch::Tensor A, torch::Tensor B,- torch::Tensor SFA, torch::Tensor SFB,- int m, int k, int l, int n_pad+ template<int ROWS_PER_BLOCK, int THREADS_PER_ROW>+ __global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)+ void gemv_wide_l2_kernel(+ 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, int N_pad) {- constexpr auto CS = CacheHint::CS;- constexpr auto CA = CacheHint::CA;+ const int row = blockIdx.x * ROWS_PER_BLOCK + threadIdx.x / THREADS_PER_ROW;+ const int batch = blockIdx.y;+ const int lane = threadIdx.x % THREADS_PER_ROW;- const int k_bytes = k / 2;- const bool b_fits_l1 = (k_bytes <= 32 * 1024);+ if (row >= M) return;- if (k >= 8192) {- if (b_fits_l1) {- dispatch<8, 32, CS, CA>(C, A, B, SFA, SFB, m, k, l, n_pad);- } else {- dispatch<8, 32, CS, CS>(C, A, B, SFA, SFB, m, k, l, n_pad);- }- } else if (k >= 4096) {- dispatch<4, 32, CS, CA>(C, A, B, SFA, SFB, m, k, l, n_pad);- } else {- dispatch<8, 16, CS, CA>(C, A, B, SFA, SFB, m, k, l, n_pad);+ const int K_bytes = K / 2;+ const int K_sf = K / 16;++ const uint8_t* a_ptr = a + (size_t)batch * M * K_bytes + row * K_bytes;+ const uint8_t* b_ptr = b + (size_t)batch * N_pad * K_bytes;+ const uint8_t* sfa_ptr = sfa + (size_t)batch * M * K_sf + row * K_sf;+ const uint8_t* sfb_ptr = sfb + (size_t)batch * N_pad * K_sf;++ float acc = 0.0f;+ const int num_sf_quads = K_sf / 4;++ #pragma unroll 2+ for (int sq = lane; sq < num_sf_quads; sq += THREADS_PER_ROW) {+ const int sf_idx = sq * 4;+ const int byte_idx = sf_idx * 8;++ uint32_t sf_a, sf_b;+ load_u32_lu(&sf_a, sfa_ptr + sf_idx);+ load_u32_b_l2(&sf_b, sfb_ptr + sf_idx);++ uint32_t a_regs[8], b_regs[8];+ asm volatile("ld.global.lu.v4.u32 {%0,%1,%2,%3}, [%4];"+ : "=r"(a_regs[0]), "=r"(a_regs[1]), "=r"(a_regs[2]), "=r"(a_regs[3])+ : "l"(a_ptr + byte_idx));+ asm volatile("ld.global.lu.v4.u32 {%0,%1,%2,%3}, [%4];"+ : "=r"(a_regs[4]), "=r"(a_regs[5]), "=r"(a_regs[6]), "=r"(a_regs[7])+ : "l"(a_ptr + byte_idx + 16));+ load_vec4_b_l2(b_regs, b_ptr + byte_idx);+ load_vec4_b_l2(b_regs + 4, b_ptr + byte_idx + 16);++ acc += dot_scaled_64fp4(a_regs, b_regs, sf_a, sf_b);}++ #pragma unroll+ for (int off = THREADS_PER_ROW / 2; off > 0; off /= 2)+ acc += __shfl_down_sync(0xffffffff, acc, off);++ if (lane == 0)+ c[(size_t)row + (size_t)batch * M] = __float2half(acc);}- '''- DIRECT_CPP = r'''- #include <torch/extension.h>- void run_direct(torch::Tensor C, torch::Tensor A, torch::Tensor B,- torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad);- '''-- # =============================================================================- # PIPELINED KERNEL MODULE (exact copy from v14)- # =============================================================================-- PIPELINED_CUDA = r'''- #include <torch/extension.h>- #include <cuda_fp16.h>- #include <cuda_fp4.h>- #include <cuda_fp8.h>-#define ASYNC_CP_16(dst, src) \asm volatile("cp.async.cg.shared.global [%0], [%1], 16;" \:: "r"((uint32_t)__cvta_generic_to_shared(dst)), "l"(src))⋯ 3 unchanged lines#define ASYNC_COMMIT() asm volatile("cp.async.commit_group;")#define ASYNC_WAIT(n) asm volatile("cp.async.wait_group %0;" :: "n"(n))- #define NUM_STAGES 2-- __device__ __forceinline__ void cvt_fp4_to_f16(half2* out, uint32_t in) {- asm volatile(- "{\n\t"- ".reg .b8 b0, b1, b2, b3;\n\t"- "mov.b32 {b0, b1, b2, b3}, %4;\n\t"- "cvt.rn.f16x2.e2m1x2 %0, b0;\n\t"- "cvt.rn.f16x2.e2m1x2 %1, b1;\n\t"- "cvt.rn.f16x2.e2m1x2 %2, b2;\n\t"- "cvt.rn.f16x2.e2m1x2 %3, b3;\n\t"- "}"- : "=r"(reinterpret_cast<int&>(out[0])),- "=r"(reinterpret_cast<int&>(out[1])),- "=r"(reinterpret_cast<int&>(out[2])),- "=r"(reinterpret_cast<int&>(out[3]))- : "r"(in)- );- }-- __device__ __forceinline__ half2 cvt_fp8_to_f16(uint16_t in) {- int out;- asm volatile("cvt.rn.f16x2.e4m3x2 %0, %1;" : "=r"(out) : "h"(in));- return *reinterpret_cast<half2*>(&out);- }-- __device__ __forceinline__ float dot_scaled_32fp4(- const uint8_t* a_smem,- const uint8_t* b_smem,- half2 scale- ) {- const uint32_t* a_data = reinterpret_cast<const uint32_t*>(a_smem);- const uint32_t* b_data = reinterpret_cast<const uint32_t*>(b_smem);-- half2 a_f16[16], b_f16[16];- #pragma unroll- for (int i = 0; i < 4; i++) {- cvt_fp4_to_f16(a_f16 + i * 4, a_data[i]);- cvt_fp4_to_f16(b_f16 + i * 4, b_data[i]);- }-- half2 acc0 = __hmul2(a_f16[0], b_f16[0]);- half2 acc1 = __hmul2(a_f16[8], b_f16[8]);-- #pragma unroll- for (int i = 1; i < 8; i++) {- acc0 = __hfma2(a_f16[i], b_f16[i], acc0);- acc1 = __hfma2(a_f16[8 + i], b_f16[8 + i], acc1);- }-- half sum0 = __hadd(acc0.x, acc0.y);- half sum1 = __hadd(acc1.x, acc1.y);-- float result = 0.0f;- asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"- : "+f"(result) : "h"(*(uint16_t*)&sum0), "h"(*(uint16_t*)&scale.x));- asm volatile("fma.rn.f32.f16 %0, %1, %2, %0;"- : "+f"(result) : "h"(*(uint16_t*)&sum1), "h"(*(uint16_t*)&scale.y));-- return result;- }-template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, int K_TILE>__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)void gemv_pipelined_kernel(- 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, int N_pad+ 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, int N_pad) {constexpr int K_TILE_BYTES = K_TILE / 2;constexpr int K_TILE_SF = K_TILE / 16;- __shared__ __align__(16) uint8_t sh_a[NUM_STAGES][ROWS_PER_BLOCK][K_TILE_BYTES];- __shared__ __align__(16) uint8_t sh_b[NUM_STAGES][K_TILE_BYTES];- __shared__ __align__(16) uint8_t sh_sfa[NUM_STAGES][ROWS_PER_BLOCK][K_TILE_SF];- __shared__ __align__(16) uint8_t sh_sfb[NUM_STAGES][K_TILE_SF];+ __shared__ __align__(16) uint8_t sh_a[2][ROWS_PER_BLOCK][K_TILE_BYTES];+ __shared__ __align__(16) uint8_t sh_b[2][K_TILE_BYTES];+ __shared__ __align__(16) uint8_t sh_sfa[2][ROWS_PER_BLOCK][K_TILE_SF];+ __shared__ __align__(16) uint8_t sh_sfb[2][K_TILE_SF];const int row_base = blockIdx.x * ROWS_PER_BLOCK;const int batch = blockIdx.y;const int tid = threadIdx.x;-const int warp_id = tid / THREADS_PER_ROW;const int lane = tid % THREADS_PER_ROW;-const int row = row_base + warp_id;+if (row >= M) return;const int K_bytes = K / 2;⋯ 13 unchanged linesif (warp_id < ROWS_PER_BLOCK && (row_base + warp_id) < M) {const uint8_t* a_row = a + a_base + (row_base + warp_id) * K_bytes + k_off_bytes;- for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16) {- if (i + 16 <= bytes_this) {- ASYNC_CP_16(&sh_a[stage][warp_id][i], a_row + i);- }- }+ for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16)+ if (i + 16 <= bytes_this) ASYNC_CP_16(&sh_a[stage][warp_id][i], a_row + i);+const uint8_t* sfa_row = sfa + sfa_base + (row_base + warp_id) * K_sf + k_off_sf;- for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4) {- if (i + 4 <= sf_this) {- ASYNC_CP_4(&sh_sfa[stage][warp_id][i], sfa_row + i);- }- }+ for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4)+ if (i + 4 <= sf_this) ASYNC_CP_4(&sh_sfa[stage][warp_id][i], sfa_row + i);}if (warp_id == 0) {const uint8_t* b_ptr = b + b_base + k_off_bytes;- for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16) {- if (i + 16 <= bytes_this) {- ASYNC_CP_16(&sh_b[stage][i], b_ptr + i);- }- }+ for (int i = lane * 16; i < bytes_this; i += THREADS_PER_ROW * 16)+ if (i + 16 <= bytes_this) ASYNC_CP_16(&sh_b[stage][i], b_ptr + i);+const uint8_t* sfb_ptr = sfb + sfb_base + k_off_sf;- for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4) {- if (i + 4 <= sf_this) {- ASYNC_CP_4(&sh_sfb[stage][i], sfb_ptr + i);- }- }+ for (int i = lane * 4; i < sf_this; i += THREADS_PER_ROW * 4)+ if (i + 4 <= sf_this) ASYNC_CP_4(&sh_sfb[stage][i], sfb_ptr + i);}ASYNC_COMMIT();};⋯ 5 unchanged linesfloat acc = 0.0f;for (int tile = 0; tile < num_tiles; tile++) {- const int stage = tile % NUM_STAGES;- const int next_stage = (tile + 1) % NUM_STAGES;+ const int stage = tile & 1;- if (tile + 1 < num_tiles) {- issue_tile(tile + 1, next_stage);- }+ if (tile + 1 < num_tiles)+ issue_tile(tile + 1, (tile + 1) & 1);const int k_off_sf = tile * K_TILE_SF;const int sf_this = min(K_TILE_SF, K_sf - k_off_sf);⋯ 4 unchanged linesconst int sf_idx = sp * 2;const int byte_idx = sf_idx * 8;- uint16_t sfa_packed = *reinterpret_cast<const uint16_t*>(&sh_sfa[stage][warp_id][sf_idx]);- uint16_t sfb_packed = *reinterpret_cast<const uint16_t*>(&sh_sfb[stage][sf_idx]);- half2 scale = __hmul2(cvt_fp8_to_f16(sfa_packed), cvt_fp8_to_f16(sfb_packed));+ uint16_t sf_a = *reinterpret_cast<const uint16_t*>(&sh_sfa[stage][warp_id][sf_idx]);+ uint16_t sf_b = *reinterpret_cast<const uint16_t*>(&sh_sfb[stage][sf_idx]);- acc += dot_scaled_32fp4(- &sh_a[stage][warp_id][byte_idx],- &sh_b[stage][byte_idx],- scale- );+ const uint32_t* ap = reinterpret_cast<const uint32_t*>(&sh_a[stage][warp_id][byte_idx]);+ const uint32_t* bp = reinterpret_cast<const uint32_t*>(&sh_b[stage][byte_idx]);++ acc += dot_scaled_32fp4(ap, bp, sf_a, sf_b);}if (tile + 1 < num_tiles) {⋯ 3 unchanged lines}#pragma unroll- for (int offset = THREADS_PER_ROW / 2; offset > 0; offset /= 2) {- acc += __shfl_down_sync(0xffffffff, acc, offset);- }+ for (int off = THREADS_PER_ROW / 2; off > 0; off /= 2)+ acc += __shfl_down_sync(0xffffffff, acc, off);- if (lane == 0) {+ if (lane == 0)c[(size_t)row + (size_t)batch * M] = __float2half(acc);- }}- template<int ROWS, int THREADS, int K_TILE>- void dispatch(torch::Tensor& c, torch::Tensor& a, torch::Tensor& b,- torch::Tensor& sfa, torch::Tensor& sfb, int m, int k, int l, int n) {- dim3 grid((m + ROWS - 1) / ROWS, l);- dim3 block(ROWS * THREADS);- gemv_pipelined_kernel<ROWS, THREADS, K_TILE><<<grid, block>>>(- a.data_ptr<uint8_t>(), b.data_ptr<uint8_t>(),- sfa.data_ptr<uint8_t>(), sfb.data_ptr<uint8_t>(),- reinterpret_cast<half*>(c.data_ptr<at::Half>()), m, k, l, n);- }-- void run_pipelined(+ void run_gemv(torch::Tensor C, torch::Tensor A, torch::Tensor B,torch::Tensor SFA, torch::Tensor SFB,- int m, int k, int l, int n_pad, int k_tile_hint+ int m, int k, int l, int n_pad) {- if (k >= 8192) {- // Try different K_TILE values for large K- // k_tile_hint: 0=2048, 1=1024, 2=4096, 3=8192- switch (k_tile_hint) {- case 1: dispatch<4, 32, 1024>(C, A, B, SFA, SFB, m, k, l, n_pad); break;- case 2: dispatch<4, 32, 4096>(C, A, B, SFA, SFB, m, k, l, n_pad); break;- case 3: dispatch<4, 32, 8192>(C, A, B, SFA, SFB, m, k, l, n_pad); break;- default: dispatch<4, 32, 2048>(C, A, B, SFA, SFB, m, k, l, n_pad); break;+ auto a_ptr = A.data_ptr<uint8_t>();+ auto b_ptr = B.data_ptr<uint8_t>();+ auto sfa_ptr = SFA.data_ptr<uint8_t>();+ auto sfb_ptr = SFB.data_ptr<uint8_t>();+ auto c_ptr = reinterpret_cast<half*>(C.data_ptr<at::Half>());++ bool use_pipelined = (k >= 8192 && l <= 2) || (k <= 4096 && l >= 4);++ if (use_pipelined) {+ dim3 grid((m + 3) / 4, l);+ dim3 block(128);++ if (k >= 8192) {+ gemv_pipelined_kernel<4, 32, 4096><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ } else {+ gemv_pipelined_kernel<4, 32, 2048><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);}- } else if (k >= 4096) {- dispatch<4, 32, 1024>(C, A, B, SFA, SFB, m, k, l, n_pad);- } else {- dispatch<8, 16, 2048>(C, A, B, SFA, SFB, m, k, l, n_pad);+ }+ else if (l >= 8) {+ dim3 grid((m + 7) / 8, l);+ dim3 block(128);+ gemv_wide_l2_kernel<8, 16><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);}+ else if (l >= 4 && k >= 4096) {+ dim3 grid((m + 7) / 8, l);+ dim3 block(128);+ gemv_direct_kernel<8, 16, LU_CA><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ }+ else {+ dim3 grid((m + 3) / 4, l);+ dim3 block(128);+ gemv_direct_kernel<4, 32, CS_CA><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ }}'''- PIPELINED_CPP = r'''+ CPP_SRC = r'''#include <torch/extension.h>- void run_pipelined(torch::Tensor C, torch::Tensor A, torch::Tensor B,- torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad, int k_tile_hint);+ void run_gemv(torch::Tensor C, torch::Tensor A, torch::Tensor B,+ torch::Tensor SFA, torch::Tensor SFB, int m, int k, int l, int n_pad);'''- # =============================================================================- # Module loading- # =============================================================================+ _module = None- # =============================================================================- # Configuration- # =============================================================================- # K_TILE options for k >= 8192:- # 0 = 2048 (default)- # 1 = 1024- # 2 = 4096- # 3 = 8192- K_TILE_HINT = 0 # <-- Change this to test different K_TILE values-- _direct_module = None- _pipelined_module = None-- def _load_direct():- global _direct_module- if _direct_module is None:- _direct_module = load_inline(- name='nvfp4_gemv_direct_v16',- cpp_sources=DIRECT_CPP,- cuda_sources=DIRECT_CUDA,- functions=['run_direct'],+ def _load_module():+ global _module+ if _module is None:+ _module = load_inline(+ name='nvfp4_gemv_v24',+ cpp_sources=CPP_SRC,+ cuda_sources=CUDA_SRC,+ functions=['run_gemv'],extra_cuda_cflags=['-O3', '--use_fast_math', '-std=c++17','-gencode=arch=compute_100a,code=sm_100a',],verbose=False,)- return _direct_module+ return _module- def _load_pipelined():- global _pipelined_module- if _pipelined_module is None:- _pipelined_module = load_inline(- name='nvfp4_gemv_pipelined_v16b', # Updated to force recompile- cpp_sources=PIPELINED_CPP,- cuda_sources=PIPELINED_CUDA,- functions=['run_pipelined'],- extra_cuda_cflags=[- '-O3', '--use_fast_math', '-std=c++17',- '-gencode=arch=compute_100a,code=sm_100a',- ],- verbose=False,- )- return _pipelined_module-def custom_kernel(data: input_t) -> output_t:a, b, sfa, sfb, _, _, c = datam, k_packed, l = a.shape⋯ 6 unchanged linessfb = sfb.to(a.device)c_out = c.squeeze(1)- a_u8 = a.view(torch.uint8)- b_u8 = b.view(torch.uint8)- sfa_u8 = sfa.view(torch.uint8)- sfb_u8 = sfb.view(torch.uint8)-- # Selection heuristic based on benchmarks:- # - k=16384, l=1: pipelined faster (24.6 vs 30.8)- # - k=7168, l=8: direct faster (38.5 vs 51.0)- # - k=2048, l=4: same (16.4)-- use_pipelined = (k >= 8192 and l <= 2) or (k <= 4096 and l >= 4)-- if use_pipelined:- _load_pipelined().run_pipelined(c_out, a_u8, b_u8, sfa_u8, sfb_u8, m, k, l, n_pad, 2)- else:- _load_direct().run_direct(c_out, a_u8, b_u8, sfa_u8, sfb_u8, m, k, l, n_pad)-+ _load_module().run_gemv(+ c_out,+ a.view(torch.uint8),+ b.view(torch.uint8),+ sfa.view(torch.uint8),+ sfb.view(torch.uint8),+ m, k, l, n_pad+ )return c
scrolls · 771 diff lines total
Best evidence level for this revision: reported
JSON