submission 114783
agokrani · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 747 lines, June 9 Researcher Reciprocity License v1.0.
nvfp4_gemv_v17_swizzled.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-114783?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:ff6228eb62074071cf7ceb58f03c17ed732c553e787d685cc04ced30cbea4d5f
license declaredunknown
license concludedunknown
authorsagokrani
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
- "Uncoalesced shared accesses" in cp_async_cg_16 affects ALL shapesfp4
NVFP4 Block-Scaled GEMV with XOR Swizzled Shared Memory Access.shared-memory
extern __shared__ uint8_t smem_raw[];vector-width = half2
half2 fallback = __float22half2_rn(make_float2(0.0f, 0.0f));Kernel source
nvfp4_gemv_v17_swizzled.py747 lines
"""
NVFP4 Block-Scaled GEMV with XOR Swizzled Shared Memory Access.
VERSION: v17_swizzled
Building on v15, this version eliminates shared memory bank conflicts
using XOR-based swizzling for A/B matrices and padding for scale arrays.
Key insight from profiling:
- "Uncoalesced shared accesses" in cp_async_cg_16 affects ALL shapes
- 16-byte stride creates 4-way bank conflicts (threads 0,8,16,24 conflict)
- For L=1, this cascades into occupancy drops and scoreboard stalls
Fix:
- XOR swizzle: index ^ (index >> 3) spreads threads across bank groups
- Padding: +4 floats per row in SFA/SFB eliminates row-based conflicts
Bank conflict analysis (before fix):
Thread 0: banks 0-3
Thread 8: banks 0-3 <- CONFLICT!
Thread 16: banks 0-3 <- CONFLICT!
Thread 24: banks 0-3 <- CONFLICT!
After XOR swizzle (thread_id ^ (thread_id >> 3)):
Thread 0 -> swizzled 0: banks 0-3
Thread 8 -> swizzled 9: banks 4-7 <- NO CONFLICT!
Thread 16 -> swizzled 18: banks 8-11 <- NO CONFLICT!
Thread 24 -> swizzled 27: banks 12-15 <- NO CONFLICT!
"""
import torch
import os
from pathlib import Path
from torch.utils.cpp_extension import load_inline
KERNEL_SOURCE = r'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_pipeline.h>
// ============================================================================
// Configuration Constants
// ============================================================================
// Standard config (for batched shapes with L >= 2)
constexpr int THREADS_PER_BLOCK = 128;
constexpr int THREADS_PER_ROW_STD = 16;
constexpr int ROWS_PER_BLOCK_STD = THREADS_PER_BLOCK / THREADS_PER_ROW_STD; // 8
constexpr int ELEMENTS_PER_ACCESS = 32;
constexpr int BYTES_PER_ACCESS = 16;
constexpr int SF_VEC_SIZE = 16;
// Wide-K config (for L=1 large-K shapes - better tail utilization)
constexpr int THREADS_PER_ROW_WIDE = 32;
constexpr int ROWS_PER_BLOCK_WIDE = THREADS_PER_BLOCK / THREADS_PER_ROW_WIDE; // 4
// Small K config (for K < 512)
constexpr int THREADS_PER_ROW_SMALL = 4;
constexpr int ROWS_PER_BLOCK_SMALL = THREADS_PER_BLOCK / THREADS_PER_ROW_SMALL; // 32
// Padding for scale arrays to avoid bank conflicts between rows
// Each row gets +4 floats (16 bytes) padding to shift bank alignment
constexpr int SFA_PAD = 4;
// ============================================================================
// Helper Functions
// ============================================================================
__device__ __forceinline__ void cp_async_cg_16(void* dst, const void* src) {
uint32_t dst_smem = static_cast<uint32_t>(__cvta_generic_to_shared(dst));
asm volatile(
"cp.async.cg.shared.global.L2::128B [%0], [%1], 16;"
: : "r"(dst_smem), "l"(src)
: "memory");
}
__device__ __forceinline__ void commit_group() { asm volatile("cp.async.commit_group;"); }
__device__ __forceinline__ void wait_group_0() { asm volatile("cp.async.wait_group 0;"); }
__device__ __forceinline__ uint32_t decode_fp4_hw(uint32_t packed_byte) {
uint32_t result;
#if defined(__CUDA_ARCH__) && __CUDA_ARCH__ >= 100
asm volatile(
"{"
".reg .b8 packed_b8;"
"cvt.u8.u32 packed_b8, %1;"
"cvt.rn.f16x2.e2m1x2 %0, packed_b8;"
"}"
: "=r"(result)
: "r"(packed_byte)
);
#else
half2 fallback = __float22half2_rn(make_float2(0.0f, 0.0f));
result = *reinterpret_cast<uint32_t*>(&fallback);
#endif
return result;
}
__device__ __forceinline__ half2 to_half2(uint32_t x) {
return *reinterpret_cast<half2*>(&x);
}
// XOR swizzle to eliminate bank conflicts
// Maps threads 0,8,16,24 to different bank groups instead of all hitting banks 0-3
__device__ __forceinline__ int xor_swizzle(int idx) {
return idx ^ (idx >> 3);
}
// ============================================================================
// Shared Memory Structures with Padding
// ============================================================================
// Standard config shared memory
// A and B use XOR swizzle for access, no structural change needed
// SFA/SFB get padding to avoid row conflicts
struct alignas(128) SmemBufferStdSwizzled {
uint8_t A[THREADS_PER_BLOCK * BYTES_PER_ACCESS]; // 2048 bytes
uint8_t B[THREADS_PER_ROW_STD * BYTES_PER_ACCESS]; // 256 bytes
float SFA[ROWS_PER_BLOCK_STD][THREADS_PER_ROW_STD * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD]; // 8 x 36
float SFB[THREADS_PER_ROW_STD * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD]; // 36
};
// Wide-K config shared memory
struct alignas(128) SmemBufferWideKSwizzled {
uint8_t A[THREADS_PER_BLOCK * BYTES_PER_ACCESS]; // 2048 bytes
uint8_t B[THREADS_PER_ROW_WIDE * BYTES_PER_ACCESS]; // 512 bytes
float SFA[ROWS_PER_BLOCK_WIDE][THREADS_PER_ROW_WIDE * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD]; // 4 x 68
float SFB[THREADS_PER_ROW_WIDE * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD]; // 68
};
// Small K config shared memory
struct alignas(128) SmemBufferSmallKSwizzled {
uint8_t A[THREADS_PER_BLOCK * BYTES_PER_ACCESS];
uint8_t B[THREADS_PER_ROW_SMALL * BYTES_PER_ACCESS];
float SFA[ROWS_PER_BLOCK_SMALL][THREADS_PER_ROW_SMALL * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD]; // 32 x 12
float SFB[THREADS_PER_ROW_SMALL * ELEMENTS_PER_ACCESS / SF_VEC_SIZE + SFA_PAD]; // 12
};
// ============================================================================
// Small K Kernel (K < 512) with XOR Swizzle
// ============================================================================
__global__ void nvfp4_gemv_kernel_v17_small_k(
const uint8_t* __restrict__ A_packed,
const uint8_t* __restrict__ B_packed,
const float* __restrict__ SFA,
const float* __restrict__ SFB,
half* __restrict__ C,
const int M,
const int K,
const int L,
const int num_m_tiles,
const int stride_A_m,
const int stride_A_l,
const int stride_B_l,
const int stride_SFA_0,
const int stride_SFA_1,
const int stride_SFA_2,
const int stride_SFA_3,
const int stride_SFA_4,
const int stride_SFA_5,
const int stride_SFB_3,
const int stride_SFB_4,
const int stride_SFB_5,
const int stride_C_m,
const int stride_C_l
) {
extern __shared__ uint8_t smem_raw[];
SmemBufferSmallKSwizzled* smem = reinterpret_cast<SmemBufferSmallKSwizzled*>(smem_raw);
const int thread_id = threadIdx.x;
const int k_thread = thread_id % THREADS_PER_ROW_SMALL;
const int m_local = thread_id / THREADS_PER_ROW_SMALL;
// XOR swizzle indices for bank conflict avoidance
const int sw_tid = xor_swizzle(thread_id);
const int sw_k = xor_swizzle(k_thread);
const int batch_idx = (gridDim.z > 1) ? blockIdx.z : (blockIdx.x / num_m_tiles);
const int m_tile_idx = (gridDim.z > 1) ? blockIdx.x : (blockIdx.x % num_m_tiles);
if (m_tile_idx >= num_m_tiles) return;
const int m_base = m_tile_idx * ROWS_PER_BLOCK_SMALL;
const int m_row = m_base + m_local;
if (m_row >= M) return;
const uint8_t* ptr_A = A_packed + batch_idx * stride_A_l;
const uint8_t* ptr_B = B_packed + batch_idx * stride_B_l;
half* ptr_C = C + batch_idx * stride_C_l + m_row * stride_C_m;
const int mm32 = m_row % 32;
const int mm4 = (m_row / 32) % 4;
const int mm = m_row / 128;
const int base_sfa_offset = mm32 * stride_SFA_0 + mm4 * stride_SFA_1 +
mm * stride_SFA_2 + batch_idx * stride_SFA_5;
const int base_sfb_offset = batch_idx * stride_SFB_5;
const int tile_k = THREADS_PER_ROW_SMALL * ELEMENTS_PER_ACCESS; // 128
const int tile_k_bytes = tile_k / 2;
const int total_k_tiles = K / tile_k;
const int scales_per_tile = tile_k / SF_VEC_SIZE;
float accum = 0.0f;
for (int tile_k_idx = 0; tile_k_idx < total_k_tiles; tile_k_idx++) {
const int k_byte_offset = tile_k_idx * tile_k_bytes;
const int scale_k_base = tile_k_idx * scales_per_tile;
// Load A with XOR swizzle
const int a_offset = m_row * stride_A_m + k_byte_offset + k_thread * 16;
cp_async_cg_16(&smem->A[sw_tid * 16], ptr_A + a_offset);
// Load B with XOR swizzle
if (m_local == 0) {
cp_async_cg_16(&smem->B[sw_k * 16], ptr_B + k_byte_offset + k_thread * 16);
}
// Load SFA (with padding in array)
if (k_thread < 2) {
int vec_idx = k_thread * 4;
int scale_idx = scale_k_base + vec_idx;
int kk4 = scale_idx % 4;
int kk = scale_idx / 4;
int off = base_sfa_offset + kk4 * stride_SFA_3 + kk * stride_SFA_4;
cp_async_cg_16(&smem->SFA[m_local][vec_idx], SFA + off);
}
// Load SFB (with padding in array)
if (m_local == 0 && k_thread < 2) {
int vec_idx = k_thread * 4;
int scale_idx = scale_k_base + vec_idx;
int kk4 = scale_idx % 4;
int kk = scale_idx / 4;
int off = base_sfb_offset + kk4 * stride_SFB_3 + kk * stride_SFB_4;
cp_async_cg_16(&smem->SFB[vec_idx], SFB + off);
}
commit_group();
wait_group_0();
__syncthreads();
// Read with XOR swizzle (same indices as write)
const uint4* vec_a = reinterpret_cast<const uint4*>(&smem->A[sw_tid * 16]);
const uint4* vec_b = reinterpret_cast<const uint4*>(&smem->B[sw_k * 16]);
uint4 va = *vec_a;
uint4 vb = *vec_b;
half2 acc_lo = __float2half2_rn(0.0f);
half2 acc_hi = __float2half2_rn(0.0f);
#define PROC(Ra, Rb, S) { \
half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
acc_lo = __hfma2(ha, hb, acc_lo); \
}
PROC(va.x, vb.x, 0); PROC(va.x, vb.x, 8); PROC(va.x, vb.x, 16); PROC(va.x, vb.x, 24);
PROC(va.y, vb.y, 0); PROC(va.y, vb.y, 8); PROC(va.y, vb.y, 16); PROC(va.y, vb.y, 24);
#undef PROC
#define PROC_HI(Ra, Rb, S) { \
half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
acc_hi = __hfma2(ha, hb, acc_hi); \
}
PROC_HI(va.z, vb.z, 0); PROC_HI(va.z, vb.z, 8); PROC_HI(va.z, vb.z, 16); PROC_HI(va.z, vb.z, 24);
PROC_HI(va.w, vb.w, 0); PROC_HI(va.w, vb.w, 8); PROC_HI(va.w, vb.w, 16); PROC_HI(va.w, vb.w, 24);
#undef PROC_HI
float sfa0 = smem->SFA[m_local][k_thread * 2];
float sfa1 = smem->SFA[m_local][k_thread * 2 + 1];
float sfb0 = smem->SFB[k_thread * 2];
float sfb1 = smem->SFB[k_thread * 2 + 1];
float sum_lo = __half2float(acc_lo.x) + __half2float(acc_lo.y);
float sum_hi = __half2float(acc_hi.x) + __half2float(acc_hi.y);
accum += sum_lo * sfa0 * sfb0 + sum_hi * sfa1 * sfb1;
__syncthreads();
}
#pragma unroll
for (int mask = THREADS_PER_ROW_SMALL / 2; mask > 0; mask >>= 1) {
accum += __shfl_xor_sync(0xFFFFFFFF, accum, mask, 32);
}
if (k_thread == 0) {
*ptr_C = __float2half(accum);
}
}
// ============================================================================
// Wide-K Kernel (K >= 8K, L=1) with XOR Swizzle
// ============================================================================
__global__ void nvfp4_gemv_kernel_v17_wide_k(
const uint8_t* __restrict__ A_packed,
const uint8_t* __restrict__ B_packed,
const float* __restrict__ SFA,
const float* __restrict__ SFB,
half* __restrict__ C,
const int M,
const int K,
const int L,
const int num_m_tiles,
const int stride_A_m,
const int stride_A_l,
const int stride_B_l,
const int stride_SFA_0,
const int stride_SFA_1,
const int stride_SFA_2,
const int stride_SFA_3,
const int stride_SFA_4,
const int stride_SFA_5,
const int stride_SFB_3,
const int stride_SFB_4,
const int stride_SFB_5,
const int stride_C_m,
const int stride_C_l
) {
extern __shared__ uint8_t smem_raw[];
SmemBufferWideKSwizzled* smem = reinterpret_cast<SmemBufferWideKSwizzled*>(smem_raw);
const int thread_id = threadIdx.x;
const int k_thread = thread_id % THREADS_PER_ROW_WIDE; // 0-31
const int m_local = thread_id / THREADS_PER_ROW_WIDE; // 0-3
// XOR swizzle indices
const int sw_tid = xor_swizzle(thread_id);
const int sw_k = xor_swizzle(k_thread);
const int m_tile_idx = blockIdx.x;
if (m_tile_idx >= num_m_tiles) return;
const int m_base = m_tile_idx * ROWS_PER_BLOCK_WIDE;
const int m_row = m_base + m_local;
if (m_row >= M) return;
const uint8_t* ptr_A = A_packed;
const uint8_t* ptr_B = B_packed;
half* ptr_C = C + m_row * stride_C_m;
const int mm32 = m_row % 32;
const int mm4 = (m_row / 32) % 4;
const int mm = m_row / 128;
const int base_sfa_offset = mm32 * stride_SFA_0 + mm4 * stride_SFA_1 + mm * stride_SFA_2;
const int base_sfb_offset = 0;
const int tile_k = THREADS_PER_ROW_WIDE * ELEMENTS_PER_ACCESS; // 1024
const int tile_k_bytes = tile_k / 2;
const int total_k_tiles = K / tile_k;
const int scales_per_tile = tile_k / SF_VEC_SIZE; // 64
float accum = 0.0f;
for (int tile_k_idx = 0; tile_k_idx < total_k_tiles; tile_k_idx++) {
const int k_byte_offset = tile_k_idx * tile_k_bytes;
const int scale_k_base = tile_k_idx * scales_per_tile;
// Load A with XOR swizzle
const int a_offset = m_row * stride_A_m + k_byte_offset + k_thread * 16;
cp_async_cg_16(&smem->A[sw_tid * 16], ptr_A + a_offset);
// Load B with XOR swizzle
if (m_local == 0) {
cp_async_cg_16(&smem->B[sw_k * 16], ptr_B + k_byte_offset + k_thread * 16);
}
// Load SFA (16 threads load 4 floats each = 64 scales)
if (k_thread < 16) {
int vec_idx = k_thread * 4;
int scale_idx = scale_k_base + vec_idx;
int kk4 = scale_idx % 4;
int kk = scale_idx / 4;
int off = base_sfa_offset + kk4 * stride_SFA_3 + kk * stride_SFA_4;
cp_async_cg_16(&smem->SFA[m_local][vec_idx], SFA + off);
}
// Load SFB
if (m_local == 0 && k_thread < 16) {
int vec_idx = k_thread * 4;
int scale_idx = scale_k_base + vec_idx;
int kk4 = scale_idx % 4;
int kk = scale_idx / 4;
int off = base_sfb_offset + kk4 * stride_SFB_3 + kk * stride_SFB_4;
cp_async_cg_16(&smem->SFB[vec_idx], SFB + off);
}
commit_group();
wait_group_0();
__syncthreads();
// Read with XOR swizzle
const uint4* vec_a = reinterpret_cast<const uint4*>(&smem->A[sw_tid * 16]);
const uint4* vec_b = reinterpret_cast<const uint4*>(&smem->B[sw_k * 16]);
uint4 va = *vec_a;
uint4 vb = *vec_b;
half2 acc_lo = __float2half2_rn(0.0f);
half2 acc_hi = __float2half2_rn(0.0f);
#define PROC(Ra, Rb, S) { \
half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
acc_lo = __hfma2(ha, hb, acc_lo); \
}
PROC(va.x, vb.x, 0); PROC(va.x, vb.x, 8); PROC(va.x, vb.x, 16); PROC(va.x, vb.x, 24);
PROC(va.y, vb.y, 0); PROC(va.y, vb.y, 8); PROC(va.y, vb.y, 16); PROC(va.y, vb.y, 24);
#undef PROC
#define PROC_HI(Ra, Rb, S) { \
half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
acc_hi = __hfma2(ha, hb, acc_hi); \
}
PROC_HI(va.z, vb.z, 0); PROC_HI(va.z, vb.z, 8); PROC_HI(va.z, vb.z, 16); PROC_HI(va.z, vb.z, 24);
PROC_HI(va.w, vb.w, 0); PROC_HI(va.w, vb.w, 8); PROC_HI(va.w, vb.w, 16); PROC_HI(va.w, vb.w, 24);
#undef PROC_HI
float sfa0 = smem->SFA[m_local][k_thread * 2];
float sfa1 = smem->SFA[m_local][k_thread * 2 + 1];
float sfb0 = smem->SFB[k_thread * 2];
float sfb1 = smem->SFB[k_thread * 2 + 1];
float sum_lo = __half2float(acc_lo.x) + __half2float(acc_lo.y);
float sum_hi = __half2float(acc_hi.x) + __half2float(acc_hi.y);
accum += sum_lo * sfa0 * sfb0 + sum_hi * sfa1 * sfb1;
__syncthreads();
}
// Warp reduction
#pragma unroll
for (int mask = THREADS_PER_ROW_WIDE / 2; mask > 0; mask >>= 1) {
accum += __shfl_xor_sync(0xFFFFFFFF, accum, mask, 32);
}
if (k_thread == 0) {
*ptr_C = __float2half(accum);
}
}
// ============================================================================
// Standard Kernel (K >= 512, L >= 2) with XOR Swizzle
// ============================================================================
__global__ void nvfp4_gemv_kernel_v17_std(
const uint8_t* __restrict__ A_packed,
const uint8_t* __restrict__ B_packed,
const float* __restrict__ SFA,
const float* __restrict__ SFB,
half* __restrict__ C,
const int M,
const int K,
const int L,
const int num_m_tiles,
const int stride_A_m,
const int stride_A_l,
const int stride_B_l,
const int stride_SFA_0,
const int stride_SFA_1,
const int stride_SFA_2,
const int stride_SFA_3,
const int stride_SFA_4,
const int stride_SFA_5,
const int stride_SFB_3,
const int stride_SFB_4,
const int stride_SFB_5,
const int stride_C_m,
const int stride_C_l
) {
extern __shared__ uint8_t smem_raw[];
SmemBufferStdSwizzled* smem = reinterpret_cast<SmemBufferStdSwizzled*>(smem_raw);
const int thread_id = threadIdx.x;
const int k_thread = thread_id % THREADS_PER_ROW_STD;
const int m_local = thread_id / THREADS_PER_ROW_STD;
// XOR swizzle indices
const int sw_tid = xor_swizzle(thread_id);
const int sw_k = xor_swizzle(k_thread);
const int m_tile_idx = blockIdx.x;
const int batch_idx = blockIdx.z;
if (m_tile_idx >= num_m_tiles) return;
const int m_base = m_tile_idx * ROWS_PER_BLOCK_STD;
const int m_row = m_base + m_local;
if (m_row >= M) return;
const uint8_t* ptr_A = A_packed + batch_idx * stride_A_l;
const uint8_t* ptr_B = B_packed + batch_idx * stride_B_l;
half* ptr_C = C + batch_idx * stride_C_l + m_row * stride_C_m;
const int mm32 = m_row % 32;
const int mm4 = (m_row / 32) % 4;
const int mm = m_row / 128;
const int base_sfa_offset = mm32 * stride_SFA_0 + mm4 * stride_SFA_1 +
mm * stride_SFA_2 + batch_idx * stride_SFA_5;
const int base_sfb_offset = batch_idx * stride_SFB_5;
const int tile_k = THREADS_PER_ROW_STD * ELEMENTS_PER_ACCESS; // 512
const int tile_k_bytes = tile_k / 2;
const int total_k_tiles = K / tile_k;
const int scales_per_tile = tile_k / SF_VEC_SIZE;
float accum = 0.0f;
for (int tile_k_idx = 0; tile_k_idx < total_k_tiles; tile_k_idx++) {
const int k_byte_offset = tile_k_idx * tile_k_bytes;
const int scale_k_base = tile_k_idx * scales_per_tile;
// Load A with XOR swizzle
const int a_offset = m_row * stride_A_m + k_byte_offset + k_thread * 16;
cp_async_cg_16(&smem->A[sw_tid * 16], ptr_A + a_offset);
// Load B with XOR swizzle
if (m_local == 0) {
cp_async_cg_16(&smem->B[sw_k * 16], ptr_B + k_byte_offset + k_thread * 16);
}
// Load SFA
if (k_thread < 8) {
int vec_idx = k_thread * 4;
int scale_idx = scale_k_base + vec_idx;
int kk4 = scale_idx % 4;
int kk = scale_idx / 4;
int off = base_sfa_offset + kk4 * stride_SFA_3 + kk * stride_SFA_4;
cp_async_cg_16(&smem->SFA[m_local][vec_idx], SFA + off);
}
// Load SFB
if (m_local == 0 && k_thread < 8) {
int vec_idx = k_thread * 4;
int scale_idx = scale_k_base + vec_idx;
int kk4 = scale_idx % 4;
int kk = scale_idx / 4;
int off = base_sfb_offset + kk4 * stride_SFB_3 + kk * stride_SFB_4;
cp_async_cg_16(&smem->SFB[vec_idx], SFB + off);
}
commit_group();
wait_group_0();
__syncthreads();
// Read with XOR swizzle
const uint4* vec_a = reinterpret_cast<const uint4*>(&smem->A[sw_tid * 16]);
const uint4* vec_b = reinterpret_cast<const uint4*>(&smem->B[sw_k * 16]);
uint4 va = *vec_a;
uint4 vb = *vec_b;
half2 acc_lo = __float2half2_rn(0.0f);
half2 acc_hi = __float2half2_rn(0.0f);
#define PROC(Ra, Rb, S) { \
half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
acc_lo = __hfma2(ha, hb, acc_lo); \
}
PROC(va.x, vb.x, 0); PROC(va.x, vb.x, 8); PROC(va.x, vb.x, 16); PROC(va.x, vb.x, 24);
PROC(va.y, vb.y, 0); PROC(va.y, vb.y, 8); PROC(va.y, vb.y, 16); PROC(va.y, vb.y, 24);
#undef PROC
#define PROC_HI(Ra, Rb, S) { \
half2 ha = to_half2(decode_fp4_hw((Ra >> S) & 0xFF)); \
half2 hb = to_half2(decode_fp4_hw((Rb >> S) & 0xFF)); \
acc_hi = __hfma2(ha, hb, acc_hi); \
}
PROC_HI(va.z, vb.z, 0); PROC_HI(va.z, vb.z, 8); PROC_HI(va.z, vb.z, 16); PROC_HI(va.z, vb.z, 24);
PROC_HI(va.w, vb.w, 0); PROC_HI(va.w, vb.w, 8); PROC_HI(va.w, vb.w, 16); PROC_HI(va.w, vb.w, 24);
#undef PROC_HI
float sfa0 = smem->SFA[m_local][k_thread * 2];
float sfa1 = smem->SFA[m_local][k_thread * 2 + 1];
float sfb0 = smem->SFB[k_thread * 2];
float sfb1 = smem->SFB[k_thread * 2 + 1];
float sum_lo = __half2float(acc_lo.x) + __half2float(acc_lo.y);
float sum_hi = __half2float(acc_hi.x) + __half2float(acc_hi.y);
accum += sum_lo * sfa0 * sfb0 + sum_hi * sfa1 * sfb1;
__syncthreads();
}
#pragma unroll
for (int mask = THREADS_PER_ROW_STD / 2; mask > 0; mask >>= 1) {
accum += __shfl_xor_sync(0xFFFFFFFF, accum, mask, 32);
}
if (k_thread == 0) {
*ptr_C = __float2half(accum);
}
}
// ============================================================================
// Launcher with Shape-Specialized Dispatch
// ============================================================================
torch::Tensor nvfp4_gemv_v17(
torch::Tensor A_packed,
torch::Tensor B_packed,
torch::Tensor SFA,
torch::Tensor SFB,
torch::Tensor C
) {
const int M = C.size(0);
const int K = A_packed.size(1) * 2;
const int L = C.size(2);
const int stride_A_m = A_packed.stride(0);
const int stride_A_l = A_packed.stride(2);
const int stride_B_l = B_packed.stride(2);
const int stride_SFA_0 = SFA.stride(0);
const int stride_SFA_1 = SFA.stride(1);
const int stride_SFA_2 = SFA.stride(2);
const int stride_SFA_3 = SFA.stride(3);
const int stride_SFA_4 = SFA.stride(4);
const int stride_SFA_5 = SFA.stride(5);
const int stride_SFB_3 = SFB.stride(3);
const int stride_SFB_4 = SFB.stride(4);
const int stride_SFB_5 = SFB.stride(5);
const int stride_C_m = C.stride(0);
const int stride_C_l = C.stride(2);
constexpr int K_SMALL_THRESHOLD = 512;
constexpr int K_WIDE_THRESHOLD = 8192;
if (K < K_SMALL_THRESHOLD) {
const int num_m_tiles = (M + ROWS_PER_BLOCK_SMALL - 1) / ROWS_PER_BLOCK_SMALL;
dim3 grid(num_m_tiles, 1, L);
dim3 block(THREADS_PER_BLOCK);
int smem_bytes = sizeof(SmemBufferSmallKSwizzled);
nvfp4_gemv_kernel_v17_small_k<<<grid, block, smem_bytes>>>(
A_packed.data_ptr<uint8_t>(),
B_packed.data_ptr<uint8_t>(),
SFA.data_ptr<float>(),
SFB.data_ptr<float>(),
reinterpret_cast<half*>(C.data_ptr<at::Half>()),
M, K, L,
num_m_tiles,
stride_A_m, stride_A_l, stride_B_l,
stride_SFA_0, stride_SFA_1, stride_SFA_2, stride_SFA_3, stride_SFA_4, stride_SFA_5,
stride_SFB_3, stride_SFB_4, stride_SFB_5,
stride_C_m, stride_C_l
);
} else if (L == 1 && K >= K_WIDE_THRESHOLD) {
const int num_m_tiles = (M + ROWS_PER_BLOCK_WIDE - 1) / ROWS_PER_BLOCK_WIDE;
dim3 grid(num_m_tiles);
dim3 block(THREADS_PER_BLOCK);
int smem_bytes = sizeof(SmemBufferWideKSwizzled);
nvfp4_gemv_kernel_v17_wide_k<<<grid, block, smem_bytes>>>(
A_packed.data_ptr<uint8_t>(),
B_packed.data_ptr<uint8_t>(),
SFA.data_ptr<float>(),
SFB.data_ptr<float>(),
reinterpret_cast<half*>(C.data_ptr<at::Half>()),
M, K, L,
num_m_tiles,
stride_A_m, stride_A_l, stride_B_l,
stride_SFA_0, stride_SFA_1, stride_SFA_2, stride_SFA_3, stride_SFA_4, stride_SFA_5,
stride_SFB_3, stride_SFB_4, stride_SFB_5,
stride_C_m, stride_C_l
);
} else {
const int num_m_tiles = (M + ROWS_PER_BLOCK_STD - 1) / ROWS_PER_BLOCK_STD;
dim3 grid(num_m_tiles, 1, L);
dim3 block(THREADS_PER_BLOCK);
int smem_bytes = sizeof(SmemBufferStdSwizzled);
nvfp4_gemv_kernel_v17_std<<<grid, block, smem_bytes>>>(
A_packed.data_ptr<uint8_t>(),
B_packed.data_ptr<uint8_t>(),
SFA.data_ptr<float>(),
SFB.data_ptr<float>(),
reinterpret_cast<half*>(C.data_ptr<at::Half>()),
M, K, L,
num_m_tiles,
stride_A_m, stride_A_l, stride_B_l,
stride_SFA_0, stride_SFA_1, stride_SFA_2, stride_SFA_3, stride_SFA_4, stride_SFA_5,
stride_SFB_3, stride_SFB_4, stride_SFB_5,
stride_C_m, stride_C_l
);
}
return C;
}
'''
CPP_SOURCE = r'''
#include <torch/extension.h>
torch::Tensor nvfp4_gemv_v17(
torch::Tensor A_packed,
torch::Tensor B_packed,
torch::Tensor SFA,
torch::Tensor SFB,
torch::Tensor C
);
'''
_extension = None
def _get_extension():
global _extension
if _extension is None:
import hashlib
code_hash = hashlib.md5(KERNEL_SOURCE.encode()).hexdigest()[:10]
print(f"Building XOR-swizzled kernel v17: nvfp4_gemv_v17_{code_hash}")
print(" Key optimization: XOR swizzle eliminates 4-way bank conflicts")
print(" Before: threads 0,8,16,24 all hit banks 0-3 (4-way conflict)")
print(" After: threads map to banks 0-3, 4-7, 8-11, 12-15 (no conflict)")
_extension = load_inline(
name=f"nvfp4_gemv_v17_{code_hash}",
cpp_sources=[CPP_SOURCE],
cuda_sources=[KERNEL_SOURCE],
functions=["nvfp4_gemv_v17"],
extra_cuda_cflags=["-O3", "--use_fast_math", "-lineinfo",
"-std=c++17", "-gencode", "arch=compute_100a,code=sm_100a"],
extra_ldflags=["-lcuda"],
with_cuda=True,
verbose=False
)
print("Compilation successful!")
return _extension
def custom_kernel(input_tuple):
a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = input_tuple
ext = _get_extension()
a_bytes = a.view(torch.uint8)
b_bytes = b.view(torch.uint8)
sfa_f32 = sfa_permuted.to(torch.float32)
sfb_f32 = sfb_permuted.to(torch.float32)
c_out = c.clone()
ext.nvfp4_gemv_v17(a_bytes, b_bytes, sfa_f32, sfb_f32, c_out)
return c_out
scrolls · 747 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Best evidence level for this revision: reported
JSON