submission 116515
macto · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 636 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-116515?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:5145d8bb93147d3a677aeb9bfdcee428a0b08a574bf23c15ebb0c32ae1e33c78
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.py636 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, L2_128B };
template<CachePolicy P>
__device__ __forceinline__ void load_vec4(uint32_t* dst, const void* src) {
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 == L2_128B) {
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));
} 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));
}
}
template<CachePolicy P>
__device__ __forceinline__ void load_u16(uint16_t* dst, const void* src) {
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 == L2_128B) {
asm volatile("ld.global.L2::128B.u16 %0, [%1];" : "=h"(*dst) : "l"(src));
} else {
asm volatile("ld.global.ca.u16 %0, [%1];" : "=h"(*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_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__ 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 broadcast_scale(half2& scale0, half2& scale1, uint16_t sfa, uint16_t sfb) {
asm volatile(
"{\n"
".reg .f16x2 sa, sb, sf;\n"
".reg .f16 s0, s1;\n"
"cvt.rn.f16x2.e4m3x2 sa, %2;\n"
"cvt.rn.f16x2.e4m3x2 sb, %3;\n"
"mul.rn.f16x2 sf, sa, sb;\n"
"mov.b32 {s0, s1}, sf;\n"
"mov.b32 %0, {s0, s0};\n"
"mov.b32 %1, {s1, s1};\n"
"}"
: "=r"(reinterpret_cast<int&>(scale0)),
"=r"(reinterpret_cast<int&>(scale1))
: "h"(sfa), "h"(sfb)
);
}
__device__ __forceinline__ float dot_scaled_32fp4(
const uint32_t* a, const uint32_t* b, uint16_t sfa, uint16_t sfb
) {
half2 scale0, scale1;
broadcast_scale(scale0, scale1, sfa, sfb);
half2 a_h[16], b_h[16];
cvt_fp4x4_to_h2x4(a_h + 0, a[0]);
cvt_fp4x4_to_h2x4(a_h + 4, a[1]);
cvt_fp4x4_to_h2x4(a_h + 8, a[2]);
cvt_fp4x4_to_h2x4(a_h + 12, a[3]);
cvt_fp4x4_to_h2x4(b_h + 0, b[0]);
cvt_fp4x4_to_h2x4(b_h + 4, b[1]);
cvt_fp4x4_to_h2x4(b_h + 8, b[2]);
cvt_fp4x4_to_h2x4(b_h + 12, b[3]);
half2 acc0 = __hmul2(a_h[0], b_h[0]);
half2 acc1 = __hmul2(a_h[8], b_h[8]);
acc0 = __hfma2(a_h[1], b_h[1], acc0);
acc1 = __hfma2(a_h[9], b_h[9], acc1);
acc0 = __hfma2(a_h[2], b_h[2], acc0);
acc1 = __hfma2(a_h[10], b_h[10], acc1);
acc0 = __hfma2(a_h[3], b_h[3], acc0);
acc1 = __hfma2(a_h[11], b_h[11], acc1);
acc0 = __hfma2(a_h[4], b_h[4], acc0);
acc1 = __hfma2(a_h[12], b_h[12], acc1);
acc0 = __hfma2(a_h[5], b_h[5], acc0);
acc1 = __hfma2(a_h[13], b_h[13], acc1);
acc0 = __hfma2(a_h[6], b_h[6], acc0);
acc1 = __hfma2(a_h[14], b_h[14], acc1);
acc0 = __hfma2(a_h[7], b_h[7], acc0);
acc1 = __hfma2(a_h[15], b_h[15], acc1);
acc0 = __hmul2(acc0, scale0);
acc1 = __hmul2(acc1, scale1);
half2 sum = __hadd2(acc0, acc1);
return __half2float(__hadd(sum.x, sum.y));
}
__device__ __forceinline__ void broadcast_scale_4(
half2& scale0, half2& scale1, half2& scale2, half2& scale3,
uint32_t sfa, uint32_t sfb
) {
asm volatile(
"{\n"
".reg .b16 sfa_lo, sfa_hi, sfb_lo, sfb_hi;\n"
".reg .f16x2 sa0, sa1, sb0, sb1, sf0, sf1;\n"
".reg .f16 s0, s1, s2, s3;\n"
"mov.b32 {sfa_lo, sfa_hi}, %4;\n"
"mov.b32 {sfb_lo, sfb_hi}, %5;\n"
"cvt.rn.f16x2.e4m3x2 sa0, sfa_lo;\n"
"cvt.rn.f16x2.e4m3x2 sa1, sfa_hi;\n"
"cvt.rn.f16x2.e4m3x2 sb0, sfb_lo;\n"
"cvt.rn.f16x2.e4m3x2 sb1, sfb_hi;\n"
"mul.rn.f16x2 sf0, sa0, sb0;\n"
"mul.rn.f16x2 sf1, sa1, sb1;\n"
"mov.b32 {s0, s1}, sf0;\n"
"mov.b32 {s2, s3}, sf1;\n"
"mov.b32 %0, {s0, s0};\n"
"mov.b32 %1, {s1, s1};\n"
"mov.b32 %2, {s2, s2};\n"
"mov.b32 %3, {s3, s3};\n"
"}"
: "=r"(reinterpret_cast<int&>(scale0)),
"=r"(reinterpret_cast<int&>(scale1)),
"=r"(reinterpret_cast<int&>(scale2)),
"=r"(reinterpret_cast<int&>(scale3))
: "r"(sfa), "r"(sfb)
);
}
__device__ __forceinline__ float dot_scaled_64fp4(
const uint32_t* a, const uint32_t* b, uint32_t sfa, uint32_t sfb
) {
half2 scale0, scale1, scale2, scale3;
broadcast_scale_4(scale0, scale1, scale2, scale3, sfa, sfb);
half2 a_h[32], b_h[32];
cvt_fp4x4_to_h2x4(a_h + 0, a[0]);
cvt_fp4x4_to_h2x4(a_h + 4, a[1]);
cvt_fp4x4_to_h2x4(a_h + 8, a[2]);
cvt_fp4x4_to_h2x4(a_h + 12, a[3]);
cvt_fp4x4_to_h2x4(a_h + 16, a[4]);
cvt_fp4x4_to_h2x4(a_h + 20, a[5]);
cvt_fp4x4_to_h2x4(a_h + 24, a[6]);
cvt_fp4x4_to_h2x4(a_h + 28, a[7]);
cvt_fp4x4_to_h2x4(b_h + 0, b[0]);
cvt_fp4x4_to_h2x4(b_h + 4, b[1]);
cvt_fp4x4_to_h2x4(b_h + 8, b[2]);
cvt_fp4x4_to_h2x4(b_h + 12, b[3]);
cvt_fp4x4_to_h2x4(b_h + 16, b[4]);
cvt_fp4x4_to_h2x4(b_h + 20, b[5]);
cvt_fp4x4_to_h2x4(b_h + 24, b[6]);
cvt_fp4x4_to_h2x4(b_h + 28, b[7]);
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);
}
acc0 = __hmul2(acc0, scale0);
acc1 = __hmul2(acc1, scale1);
acc2 = __hmul2(acc2, scale2);
acc3 = __hmul2(acc3, scale3);
half2 sum01 = __hadd2(acc0, acc1);
half2 sum23 = __hadd2(acc2, acc3);
half2 sum = __hadd2(sum01, sum23);
return __half2float(__hadd(sum.x, sum.y));
}
template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CachePolicy POLICY>
__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW, 8)
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<CachePolicy P>
__device__ __forceinline__ void load_vec4_wide(uint32_t* dst, const void* src) {
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_wide_b(uint32_t* dst, const void* src) {
if constexpr (P == L2_128B) {
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));
} 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));
}
}
template<CachePolicy P>
__device__ __forceinline__ void load_u32_wide(uint32_t* dst, const void* src) {
if constexpr (P == LU_CA) {
asm volatile("ld.global.lu.u32 %0, [%1];" : "=r"(*dst) : "l"(src));
} else {
asm volatile("ld.global.cs.u32 %0, [%1];" : "=r"(*dst) : "l"(src));
}
}
template<CachePolicy P>
__device__ __forceinline__ void load_u32_wide_b(uint32_t* dst, const void* src) {
if constexpr (P == L2_128B) {
asm volatile("ld.global.L2::128B.u32 %0, [%1];" : "=r"(*dst) : "l"(src));
} else {
asm volatile("ld.global.ca.u32 %0, [%1];" : "=r"(*dst) : "l"(src));
}
}
template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CachePolicy POLICY>
__global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW, 8)
void gemv_wide_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_wide<POLICY>(&sf_a, sfa_ptr + sf_idx);
load_u32_wide_b<POLICY>(&sf_b, sfb_ptr + sf_idx);
uint32_t a_regs[8], b_regs[8];
load_vec4_wide<POLICY>(a_regs, a_ptr + byte_idx);
load_vec4_wide<POLICY>(a_regs + 4, a_ptr + byte_idx + 16);
load_vec4_wide_b<POLICY>(b_regs, b_ptr + byte_idx);
load_vec4_wide_b<POLICY>(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, 8)
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();
}
}
if constexpr (THREADS_PER_ROW <= 32) {
#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);
} else {
__shared__ float smem_reduce[ROWS_PER_BLOCK][THREADS_PER_ROW];
smem_reduce[warp_id][lane] = acc;
__syncthreads();
if (lane < 64) smem_reduce[warp_id][lane] += smem_reduce[warp_id][lane + 64];
__syncthreads();
if (lane < 32) {
float val = smem_reduce[warp_id][lane] + smem_reduce[warp_id][lane + 32];
#pragma unroll
for (int off = 16; off > 0; off /= 2)
val += __shfl_down_sync(0xffffffff, val, off);
if (lane == 0)
c[(size_t)row + (size_t)batch * M] = __float2half(val);
}
}
}
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>());
if (k == 16384 && l == 1) {
#if L1_CONFIG == 0
// Config 0: 4 rows, 32 threads/row, pipelined (default)
dim3 grid((m + 3) / 4, l);
dim3 block(128);
gemv_pipelined_kernel<4, 32, 2048><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
#elif L1_CONFIG == 1
// Config 1: 8 rows, 16 threads/row, pipelined
dim3 grid((m + 7) / 8, l);
dim3 block(128);
gemv_pipelined_kernel<8, 16, 2048><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
#elif L1_CONFIG == 2
// Config 2: 4 rows, 32 threads/row, direct CS_CA
dim3 grid((m + 3) / 4, l);
dim3 block(128);
gemv_direct_kernel<4, 32, LU_CA><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
#elif L1_CONFIG == 3
// Config 3: 1 row, 128 threads, pipelined (Sam-style)
dim3 grid(m, l);
dim3 block(128);
gemv_pipelined_kernel<1, 128, 2048><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
#else
// Default
dim3 grid((m + 3) / 4, l);
dim3 block(128);
gemv_pipelined_kernel<4, 32, 4096><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
#endif
}
else if (k == 7168 && l == 8) {
dim3 grid((m + 3) / 4, l);
dim3 block(128);
gemv_direct_kernel<4, 32, LU_CA><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
}
else if (k == 2048 && l == 4) {
dim3 grid((m + 3) / 4, l);
dim3 block(128);
gemv_wide_kernel<4, 32, LU_CA><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
}
else if (l >= 8) {
dim3 grid((m + 3) / 4, l);
dim3 block(128);
gemv_direct_kernel<4, 32, LU_CA><<<grid, block>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);
}
else if (k >= 8192) {
dim3 grid((m + 3) / 4, l);
dim3 block(128);
gemv_pipelined_kernel<4, 32, 4096><<<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
# L1_CONFIG: Configuration for L=1, K=16384 case
# 0 = 4 rows, 32 threads/row, pipelined (default)
# 1 = 8 rows, 16 threads/row, pipelined
# 2 = 4 rows, 32 threads/row, direct CS_CA
# 3 = 1 row, 128 threads, pipelined (Sam-style)
L1_CONFIG = 2
def _load_module():
global _module
if _module is None:
_module = load_inline(
name=f'nvfp4_gemv_v28_l1c{L1_CONFIG}',
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',
'-maxrregcount=48',
f'-DL1_CONFIG={L1_CONFIG}',
],
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 · 636 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 113565.
⋯ 7 unchanged lines#include <cuda_fp4.h>#include <cuda_fp8.h>- enum CachePolicy { CS_CA, LU_CA, NC_EVICT };+ enum CachePolicy { CS_CA, LU_CA, L2_128B };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) {+ 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 {⋯ 4 unchanged linestemplate<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];"+ if constexpr (P == L2_128B) {+ 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));} else {asm volatile("ld.global.ca.v4.u32 {%0,%1,%2,%3}, [%4];"⋯ 1 unchanged lines}}- __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) {+ 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));⋯ 2 unchanged linestemplate<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));+ if constexpr (P == L2_128B) {+ asm volatile("ld.global.L2::128B.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_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__ void cvt_fp8x4_to_h2x4(half2* out, uint32_t in) {asm volatile("{\n\t"⋯ 8 unchanged lines);}- __device__ __forceinline__ void cvt_fp4x4_to_h2x4(half2* out, uint32_t in) {+ __device__ __forceinline__ void broadcast_scale(half2& scale0, half2& scale1, uint16_t sfa, uint16_t sfb) {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"+ "{\n"+ ".reg .f16x2 sa, sb, sf;\n"+ ".reg .f16 s0, s1;\n"+ "cvt.rn.f16x2.e4m3x2 sa, %2;\n"+ "cvt.rn.f16x2.e4m3x2 sb, %3;\n"+ "mul.rn.f16x2 sf, sa, sb;\n"+ "mov.b32 {s0, s1}, sf;\n"+ "mov.b32 %0, {s0, s0};\n"+ "mov.b32 %1, {s1, s1};\n""}"- : "=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)+ : "=r"(reinterpret_cast<int&>(scale0)),+ "=r"(reinterpret_cast<int&>(scale1))+ : "h"(sfa), "h"(sfb));}__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 scale0, scale1;+ broadcast_scale(scale0, scale1, sfa, 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]);- }+ cvt_fp4x4_to_h2x4(a_h + 0, a[0]);+ cvt_fp4x4_to_h2x4(a_h + 4, a[1]);+ cvt_fp4x4_to_h2x4(a_h + 8, a[2]);+ cvt_fp4x4_to_h2x4(a_h + 12, a[3]);+ cvt_fp4x4_to_h2x4(b_h + 0, b[0]);+ cvt_fp4x4_to_h2x4(b_h + 4, b[1]);+ cvt_fp4x4_to_h2x4(b_h + 8, b[2]);+ cvt_fp4x4_to_h2x4(b_h + 12, b[3]);+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);- }+ acc0 = __hfma2(a_h[1], b_h[1], acc0);+ acc1 = __hfma2(a_h[9], b_h[9], acc1);+ acc0 = __hfma2(a_h[2], b_h[2], acc0);+ acc1 = __hfma2(a_h[10], b_h[10], acc1);+ acc0 = __hfma2(a_h[3], b_h[3], acc0);+ acc1 = __hfma2(a_h[11], b_h[11], acc1);+ acc0 = __hfma2(a_h[4], b_h[4], acc0);+ acc1 = __hfma2(a_h[12], b_h[12], acc1);+ acc0 = __hfma2(a_h[5], b_h[5], acc0);+ acc1 = __hfma2(a_h[13], b_h[13], acc1);+ acc0 = __hfma2(a_h[6], b_h[6], acc0);+ acc1 = __hfma2(a_h[14], b_h[14], acc1);+ acc0 = __hfma2(a_h[7], b_h[7], acc0);+ acc1 = __hfma2(a_h[15], b_h[15], acc1);- half sum0 = __hadd(acc0.x, acc0.y);- half sum1 = __hadd(acc1.x, acc1.y);+ acc0 = __hmul2(acc0, scale0);+ acc1 = __hmul2(acc1, scale1);- return __half2float(__hadd(__hmul(sum0, scale.x), __hmul(sum1, scale.y)));+ half2 sum = __hadd2(acc0, acc1);+ return __half2float(__hadd(sum.x, sum.y));}+ __device__ __forceinline__ void broadcast_scale_4(+ half2& scale0, half2& scale1, half2& scale2, half2& scale3,+ uint32_t sfa, uint32_t sfb+ ) {+ asm volatile(+ "{\n"+ ".reg .b16 sfa_lo, sfa_hi, sfb_lo, sfb_hi;\n"+ ".reg .f16x2 sa0, sa1, sb0, sb1, sf0, sf1;\n"+ ".reg .f16 s0, s1, s2, s3;\n"+ "mov.b32 {sfa_lo, sfa_hi}, %4;\n"+ "mov.b32 {sfb_lo, sfb_hi}, %5;\n"+ "cvt.rn.f16x2.e4m3x2 sa0, sfa_lo;\n"+ "cvt.rn.f16x2.e4m3x2 sa1, sfa_hi;\n"+ "cvt.rn.f16x2.e4m3x2 sb0, sfb_lo;\n"+ "cvt.rn.f16x2.e4m3x2 sb1, sfb_hi;\n"+ "mul.rn.f16x2 sf0, sa0, sb0;\n"+ "mul.rn.f16x2 sf1, sa1, sb1;\n"+ "mov.b32 {s0, s1}, sf0;\n"+ "mov.b32 {s2, s3}, sf1;\n"+ "mov.b32 %0, {s0, s0};\n"+ "mov.b32 %1, {s1, s1};\n"+ "mov.b32 %2, {s2, s2};\n"+ "mov.b32 %3, {s3, s3};\n"+ "}"+ : "=r"(reinterpret_cast<int&>(scale0)),+ "=r"(reinterpret_cast<int&>(scale1)),+ "=r"(reinterpret_cast<int&>(scale2)),+ "=r"(reinterpret_cast<int&>(scale3))+ : "r"(sfa), "r"(sfb)+ );+ }+__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, scale1, scale2, scale3;+ broadcast_scale_4(scale0, scale1, scale2, scale3, sfa, 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]);- }+ cvt_fp4x4_to_h2x4(a_h + 0, a[0]);+ cvt_fp4x4_to_h2x4(a_h + 4, a[1]);+ cvt_fp4x4_to_h2x4(a_h + 8, a[2]);+ cvt_fp4x4_to_h2x4(a_h + 12, a[3]);+ cvt_fp4x4_to_h2x4(a_h + 16, a[4]);+ cvt_fp4x4_to_h2x4(a_h + 20, a[5]);+ cvt_fp4x4_to_h2x4(a_h + 24, a[6]);+ cvt_fp4x4_to_h2x4(a_h + 28, a[7]);+ cvt_fp4x4_to_h2x4(b_h + 0, b[0]);+ cvt_fp4x4_to_h2x4(b_h + 4, b[1]);+ cvt_fp4x4_to_h2x4(b_h + 8, b[2]);+ cvt_fp4x4_to_h2x4(b_h + 12, b[3]);+ cvt_fp4x4_to_h2x4(b_h + 16, b[4]);+ cvt_fp4x4_to_h2x4(b_h + 20, b[5]);+ cvt_fp4x4_to_h2x4(b_h + 24, b[6]);+ cvt_fp4x4_to_h2x4(b_h + 28, b[7]);+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]);⋯ 7 unchanged linesacc3 = __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);+ acc0 = __hmul2(acc0, scale0);+ acc1 = __hmul2(acc1, scale1);+ acc2 = __hmul2(acc2, scale2);+ acc3 = __hmul2(acc3, scale3);- half scaled0 = __hadd(__hmul(sum0, scale0.x), __hmul(sum1, scale0.y));- half scaled1 = __hadd(__hmul(sum2, scale1.x), __hmul(sum3, scale1.y));+ half2 sum01 = __hadd2(acc0, acc1);+ half2 sum23 = __hadd2(acc2, acc3);+ half2 sum = __hadd2(sum01, sum23);- return __half2float(__hadd(scaled0, scaled1));+ return __half2float(__hadd(sum.x, sum.y));}template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CachePolicy POLICY>- __global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW)+ __global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW, 8)void gemv_direct_kernel(const uint8_t* __restrict__ a, const uint8_t* __restrict__ b,const uint8_t* __restrict__ sfa, const uint8_t* __restrict__ sfb,⋯ 40 unchanged linesc[(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(+ template<CachePolicy P>+ __device__ __forceinline__ void load_vec4_wide(uint32_t* dst, const void* src) {+ 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_wide_b(uint32_t* dst, const void* src) {+ if constexpr (P == L2_128B) {+ 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));+ } 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));+ }+ }++ template<CachePolicy P>+ __device__ __forceinline__ void load_u32_wide(uint32_t* dst, const void* src) {+ if constexpr (P == LU_CA) {+ asm volatile("ld.global.lu.u32 %0, [%1];" : "=r"(*dst) : "l"(src));+ } else {+ asm volatile("ld.global.cs.u32 %0, [%1];" : "=r"(*dst) : "l"(src));+ }+ }++ template<CachePolicy P>+ __device__ __forceinline__ void load_u32_wide_b(uint32_t* dst, const void* src) {+ if constexpr (P == L2_128B) {+ asm volatile("ld.global.L2::128B.u32 %0, [%1];" : "=r"(*dst) : "l"(src));+ } else {+ asm volatile("ld.global.ca.u32 %0, [%1];" : "=r"(*dst) : "l"(src));+ }+ }++ template<int ROWS_PER_BLOCK, int THREADS_PER_ROW, CachePolicy POLICY>+ __global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW, 8)+ void gemv_wide_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⋯ 21 unchanged linesconst 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);+ load_u32_wide<POLICY>(&sf_a, sfa_ptr + sf_idx);+ load_u32_wide_b<POLICY>(&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);+ load_vec4_wide<POLICY>(a_regs, a_ptr + byte_idx);+ load_vec4_wide<POLICY>(a_regs + 4, a_ptr + byte_idx + 16);+ load_vec4_wide_b<POLICY>(b_regs, b_ptr + byte_idx);+ load_vec4_wide_b<POLICY>(b_regs + 4, b_ptr + byte_idx + 16);acc += dot_scaled_64fp4(a_regs, b_regs, sf_a, sf_b);}⋯ 16 unchanged lines#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)+ __global__ __launch_bounds__(ROWS_PER_BLOCK * THREADS_PER_ROW, 8)void gemv_pipelined_kernel(const uint8_t* __restrict__ a, const uint8_t* __restrict__ b,const uint8_t* __restrict__ sfa, const uint8_t* __restrict__ sfb,⋯ 89 unchanged lines}}- #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);+ if constexpr (THREADS_PER_ROW <= 32) {+ #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);+ } else {+ __shared__ float smem_reduce[ROWS_PER_BLOCK][THREADS_PER_ROW];+ smem_reduce[warp_id][lane] = acc;+ __syncthreads();++ if (lane < 64) smem_reduce[warp_id][lane] += smem_reduce[warp_id][lane + 64];+ __syncthreads();+ if (lane < 32) {+ float val = smem_reduce[warp_id][lane] + smem_reduce[warp_id][lane + 32];+ #pragma unroll+ for (int off = 16; off > 0; off /= 2)+ val += __shfl_down_sync(0xffffffff, val, off);+ if (lane == 0)+ c[(size_t)row + (size_t)batch * M] = __float2half(val);+ }+ }}void run_gemv(⋯ 7 unchanged linesauto 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) {+ if (k == 16384 && l == 1) {+ #if L1_CONFIG == 0+ // Config 0: 4 rows, 32 threads/row, pipelined (default)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) {+ gemv_pipelined_kernel<4, 32, 2048><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ #elif L1_CONFIG == 1+ // Config 1: 8 rows, 16 threads/row, pipelineddim3 grid((m + 7) / 8, l);dim3 block(128);- gemv_wide_l2_kernel<8, 16><<<grid, block>>>(+ gemv_pipelined_kernel<8, 16, 2048><<<grid, block>>>(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ #elif L1_CONFIG == 2+ // Config 2: 4 rows, 32 threads/row, direct CS_CA+ dim3 grid((m + 3) / 4, l);+ dim3 block(128);+ gemv_direct_kernel<4, 32, LU_CA><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ #elif L1_CONFIG == 3+ // Config 3: 1 row, 128 threads, pipelined (Sam-style)+ dim3 grid(m, l);+ dim3 block(128);+ gemv_pipelined_kernel<1, 128, 2048><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ #else+ // Default+ dim3 grid((m + 3) / 4, l);+ dim3 block(128);+ gemv_pipelined_kernel<4, 32, 4096><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ #endif}- else if (l >= 4 && k >= 4096) {- dim3 grid((m + 7) / 8, l);+ else if (k == 7168 && l == 8) {+ dim3 grid((m + 3) / 4, l);dim3 block(128);- gemv_direct_kernel<8, 16, LU_CA><<<grid, block>>>(+ gemv_direct_kernel<4, 32, LU_CA><<<grid, block>>>(a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);- }+ }+ else if (k == 2048 && l == 4) {+ dim3 grid((m + 3) / 4, l);+ dim3 block(128);+ gemv_wide_kernel<4, 32, LU_CA><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ }+ else if (l >= 8) {+ dim3 grid((m + 3) / 4, l);+ dim3 block(128);+ gemv_direct_kernel<4, 32, LU_CA><<<grid, block>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr, c_ptr, m, k, l, n_pad);+ }+ else if (k >= 8192) {+ dim3 grid((m + 3) / 4, l);+ dim3 block(128);+ gemv_pipelined_kernel<4, 32, 4096><<<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);⋯ 11 unchanged lines_module = None+ # L1_CONFIG: Configuration for L=1, K=16384 case+ # 0 = 4 rows, 32 threads/row, pipelined (default)+ # 1 = 8 rows, 16 threads/row, pipelined+ # 2 = 4 rows, 32 threads/row, direct CS_CA+ # 3 = 1 row, 128 threads, pipelined (Sam-style)+ L1_CONFIG = 2+def _load_module():global _moduleif _module is None:_module = load_inline(- name='nvfp4_gemv_v24',+ name=f'nvfp4_gemv_v28_l1c{L1_CONFIG}',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',+ '-maxrregcount=48',+ f'-DL1_CONFIG={L1_CONFIG}',],verbose=False,)
scrolls · 524 diff lines total
Best evidence level for this revision: reported
JSON