submission 70351
snowclipsed · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 251 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-70351?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:c830d9f96e20a8e2218ad4dc654f6bb0fafcea576a692ca3211e97b939ff4f9c
license declaredunknown
license concludedunknown
authorssnowclipsed
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp4
Vectorized NVFP4 GEMV with shared memory - 8x fewer K-loop iterationsshared-memory
extern __shared__ unsigned char smem[];Kernel source
submission.py251 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
nvfp4_gemv_cuda = """
#include <cuda_fp16.h>
#include <cuda_runtime.h>
// === Helper Functions (must come before kernel) ===
__device__ __forceinline__ float dequant_fp4_e2m1(unsigned char packed_val, int idx) {
unsigned char fp4_bits = (idx == 0) ? (packed_val & 0x0F) : ((packed_val >> 4) & 0x0F);
const float lut[16] = {
0.0f, 0.5f, 1.0f, 1.5f,
2.0f, 3.0f, 4.0f, 6.0f,
-0.0f, -0.5f, -1.0f, -1.5f,
-2.0f, -3.0f, -4.0f, -6.0f
};
return lut[fp4_bits];
}
__device__ __forceinline__ float dequant_fp8_e4m3(unsigned char fp8_bits) {
int sign = (fp8_bits >> 7) & 0x1;
int exp = (fp8_bits >> 3) & 0xF;
int mant = fp8_bits & 0x7;
float val;
if (exp == 0) {
val = ldexpf(mant / 8.0f, -6);
} else if (exp == 15) {
val = 448.0f;
} else {
val = ldexpf(1.0f + mant / 8.0f, exp - 7);
}
return sign ? -val : val;
}
__device__ __forceinline__ int64_t blocked_scale_offset(
int m, int k_block, int l,
int rest_m_dim, int rest_k_dim,
int64_t s0, int64_t s1, int64_t s2, int64_t s3, int64_t s4, int64_t s5
) {
int mm = m / 128;
int mm32 = m % 32;
int mm4 = (m % 128) / 32;
int kk = k_block / 4;
int kk4 = k_block % 4;
if (mm >= rest_m_dim || kk >= rest_k_dim) return -1;
return mm32 * s0 + mm4 * s1 + mm * s2 + kk4 * s3 + kk * s4 + l * s5;
}
// === Vectorized Kernel ===
__global__ void nvfp4_gemv_vectorized(
const unsigned char* __restrict__ a,
const unsigned char* __restrict__ b,
const unsigned char* __restrict__ sfa,
const unsigned char* __restrict__ sfb,
__half* __restrict__ c,
int M, int K, int L, int B_rows,
int sfa_rest_m, int sfa_rest_k, int sfb_rest_m, int sfb_rest_k,
int64_t a_s0, int64_t a_s1, int64_t a_s2,
int64_t b_s0, int64_t b_s1, int64_t b_s2,
int64_t sfa_s0, int64_t sfa_s1, int64_t sfa_s2, int64_t sfa_s3, int64_t sfa_s4, int64_t sfa_s5,
int64_t sfb_s0, int64_t sfb_s1, int64_t sfb_s2, int64_t sfb_s3, int64_t sfb_s4, int64_t sfb_s5,
int64_t c_s0, int64_t c_s1, int64_t c_s2
) {
extern __shared__ unsigned char smem[];
unsigned char* b_shared = smem;
unsigned char* sfb_shared = smem + (K / 2);
const int warp_id = threadIdx.y;
const int lane_id = threadIdx.x;
const int m = blockIdx.x * blockDim.y + warp_id;
const int l = blockIdx.y;
const int tid = threadIdx.x + threadIdx.y * blockDim.x;
if (m >= M || l >= L) return;
const int K_bytes = K / 2;
const int K_blocks = K / 16;
const int b_m = 0;
// Load b (coalesced)
#pragma unroll 8
for (int idx = tid; idx < K_bytes; idx += blockDim.x * blockDim.y) {
b_shared[idx] = __ldg(&b[idx * b_s1 + l * b_s2]);
}
// Precompute scale offsets
const int mm = m / 128;
const int mm32 = m % 32;
const int mm4 = (m % 128) / 32;
const int64_t sfa_m_part = mm32 * sfa_s0 + mm4 * sfa_s1 + mm * sfa_s2 + l * sfa_s5;
const int64_t sfb_l_offset = l * sfb_s5;
#pragma unroll 8
for (int idx = tid; idx < K_blocks; idx += blockDim.x * blockDim.y) {
const int kk = idx / 4;
const int kk4 = idx % 4;
const int64_t sfb_offset = b_m * sfb_s0 + kk4 * sfb_s3 + kk * sfb_s4 + sfb_l_offset;
sfb_shared[idx] = __ldg(&sfb[sfb_offset]);
}
__syncthreads();
// === VECTORIZED: Process 8 bytes per iteration ===
float thread_acc = 0.0f;
#pragma unroll 1
for (int k_byte = lane_id * 8; k_byte < K_bytes; k_byte += warpSize * 8) {
const int64_t a_offset = m * a_s0 + k_byte * a_s1 + l * a_s2;
// Load 8 bytes (64 bits) at once
uint2 a_vec = make_uint2(0, 0);
if (k_byte + 8 <= K_bytes) {
a_vec = *reinterpret_cast<const uint2*>(&a[a_offset]);
} else {
a_vec.x = (k_byte < K_bytes) ? *reinterpret_cast<const uint*>(&a[a_offset]) : 0;
a_vec.y = (k_byte + 4 < K_bytes) ? *reinterpret_cast<const uint*>(&a[a_offset + 4]) : 0;
}
// Process 16 FP4 values
#pragma unroll 8
for (int i = 0; i < 8; i++) {
const int byte_idx = k_byte + i;
if (byte_idx >= K_bytes) break;
const int k = byte_idx * 2;
const int k_block = byte_idx >> 3;
const unsigned char a_val = reinterpret_cast<unsigned char*>(&a_vec)[i];
const unsigned char b_val = b_shared[byte_idx];
const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0);
const float a_fp4_1 = (k + 1 < K) ? dequant_fp4_e2m1(a_val, 1) : 0.0f;
const float b_fp4_0 = dequant_fp4_e2m1(b_val, 0);
const float b_fp4_1 = dequant_fp4_e2m1(b_val, 1);
const int kk = k_block / 4;
const int kk4 = k_block % 4;
const int64_t sfa_offset = sfa_m_part + kk4 * sfa_s3 + kk * sfa_s4;
// Bounds check (computational equivalence)
if (sfa_offset >= 0) {
const float scale_a = dequant_fp8_e4m3(__ldg(&sfa[sfa_offset]));
const float scale_b = dequant_fp8_e4m3(sfb_shared[k_block]);
const float a_scaled_0 = a_fp4_0 * scale_a;
const float a_scaled_1 = a_fp4_1 * scale_a;
const float b_scaled_0 = b_fp4_0 * scale_b;
const float b_scaled_1 = b_fp4_1 * scale_b;
thread_acc += a_scaled_0 * b_scaled_0;
if (k + 1 < K) thread_acc += a_scaled_1 * b_scaled_1;
}
}
}
// Warp reduction
#pragma unroll
for (int offset = 16; offset > 0; offset >>= 1) {
thread_acc += __shfl_down_sync(0xFFFFFFFF, thread_acc, offset);
}
if (lane_id == 0) {
const int64_t c_offset = m * c_s0 + l * c_s2;
c[c_offset] = __float2half(thread_acc);
}
}
// === PyTorch wrapper (matches original signature) ===
torch::Tensor nvfp4_gemv(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa_permuted,
torch::Tensor sfb_permuted,
torch::Tensor c
) {
TORCH_CHECK(a.device().is_cuda(), "tensors must be CUDA");
TORCH_CHECK(sfa_permuted.dim() == 6, "sfa_permuted must be 6D");
TORCH_CHECK(sfb_permuted.dim() == 6, "sfb_permuted must be 6D");
unsigned char* a_ptr = reinterpret_cast<unsigned char*>(a.data_ptr());
unsigned char* b_ptr = reinterpret_cast<unsigned char*>(b.data_ptr());
unsigned char* sfa_ptr = reinterpret_cast<unsigned char*>(sfa_permuted.data_ptr());
unsigned char* sfb_ptr = reinterpret_cast<unsigned char*>(sfb_permuted.data_ptr());
int M = a.size(0);
int K_bytes = a.size(1);
int L = a.size(2);
int K = K_bytes * 2;
int B_rows = b.size(0);
int sfa_dim2 = sfa_permuted.size(2);
int sfa_dim4 = sfa_permuted.size(4);
int sfb_dim2 = sfb_permuted.size(2);
int sfb_dim4 = sfb_permuted.size(4);
dim3 block(32, 16);
dim3 grid((M + 15) / 16, L);
size_t smem_size = K_bytes + K / 16;
nvfp4_gemv_vectorized<<<grid, block, smem_size>>>(
a_ptr, b_ptr, sfa_ptr, sfb_ptr,
reinterpret_cast<__half*>(c.data_ptr<at::Half>()),
M, K, L, B_rows,
sfa_dim2, sfa_dim4, sfb_dim2, sfb_dim4,
a.stride(0), a.stride(1), a.stride(2),
b.stride(0), b.stride(1), b.stride(2),
sfa_permuted.stride(0), sfa_permuted.stride(1), sfa_permuted.stride(2),
sfa_permuted.stride(3), sfa_permuted.stride(4), sfa_permuted.stride(5),
sfb_permuted.stride(0), sfb_permuted.stride(1), sfb_permuted.stride(2),
sfb_permuted.stride(3), sfb_permuted.stride(4), sfb_permuted.stride(5),
c.stride(0), c.stride(1), c.stride(2)
);
cudaError_t err = cudaGetLastError();
TORCH_CHECK(err == cudaSuccess, "CUDA error: ", cudaGetErrorString(err));
return c;
}
"""
nvfp4_gemv_cpp = """
#include <torch/extension.h>
torch::Tensor nvfp4_gemv(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa,
torch::Tensor sfb,
torch::Tensor c
);
"""
nvfp4_module = load_inline(
name='nvfp4_gemv',
cpp_sources=nvfp4_gemv_cpp,
cuda_sources=nvfp4_gemv_cuda,
functions=['nvfp4_gemv'],
verbose=False,
extra_cuda_cflags=['-O3', '--use_fast_math', '-arch=sm_100']
)
def custom_kernel(data: input_t) -> output_t: # type: ignore
"""
Vectorized NVFP4 GEMV with shared memory - 8x fewer K-loop iterations
"""
a, b, _, _, sfa_permuted, sfb_permuted, c = data
return nvfp4_module.nvfp4_gemv(a, b, sfa_permuted, sfb_permuted, c)scrolls · 251 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 69575.
- from task import input_t, output_timport torch- import triton- import triton.language as tl+ from torch.utils.cpp_extension import load_inline+ from task import input_t, output_t- # =============================================================================- # OPTIMIZED SCALE FACTOR BLOCKING - From your original kernel- # =============================================================================+ nvfp4_gemv_cuda = """+ #include <cuda_fp16.h>+ #include <cuda_runtime.h>- def ceil_div(a, b):- return (a + b - 1) // b+ // === Helper Functions (must come before kernel) ===- @triton.jit- def blocked_transform_kernel(- inp, out, M, K, L,- s_im, s_ik, s_il, s_ol, s_oe,- BLK: tl.constexpr- ):- """Optimized blocking transformation for scale factors"""- pid_l = tl.program_id(0)- pid_b = tl.program_id(1)+ __device__ __forceinline__ float dequant_fp4_e2m1(unsigned char packed_val, int idx) {+ unsigned char fp4_bits = (idx == 0) ? (packed_val & 0x0F) : ((packed_val >> 4) & 0x0F);+ const float lut[16] = {+ 0.0f, 0.5f, 1.0f, 1.5f,+ 2.0f, 3.0f, 4.0f, 6.0f,+ -0.0f, -0.5f, -1.0f, -1.5f,+ -2.0f, -3.0f, -4.0f, -6.0f+ };+ return lut[fp4_bits];+ }++ __device__ __forceinline__ float dequant_fp8_e4m3(unsigned char fp8_bits) {+ int sign = (fp8_bits >> 7) & 0x1;+ int exp = (fp8_bits >> 3) & 0xF;+ int mant = fp8_bits & 0x7;- mk = M * K- offs = pid_b * BLK + tl.arange(0, BLK)- mask = offs < mk+ float val;+ if (exp == 0) {+ val = ldexpf(mant / 8.0f, -6);+ } else if (exp == 15) {+ val = 448.0f;+ } else {+ val = ldexpf(1.0f + mant / 8.0f, exp - 7);+ }+ return sign ? -val : val;+ }++ __device__ __forceinline__ int64_t blocked_scale_offset(+ int m, int k_block, int l,+ int rest_m_dim, int rest_k_dim,+ int64_t s0, int64_t s1, int64_t s2, int64_t s3, int64_t s4, int64_t s5+ ) {+ int mm = m / 128;+ int mm32 = m % 32;+ int mm4 = (m % 128) / 32;+ int kk = k_block / 4;+ int kk4 = k_block % 4;- i = offs // K- j = offs % K+ if (mm >= rest_m_dim || kk >= rest_k_dim) return -1;- nrb = (M + 127) // 128- ncb = (K + 3) // 4+ return mm32 * s0 + mm4 * s1 + mm * s2 + kk4 * s3 + kk * s4 + l * s5;+ }++ // === Vectorized Kernel ===++ __global__ void nvfp4_gemv_vectorized(+ const unsigned char* __restrict__ a,+ const unsigned char* __restrict__ b,+ const unsigned char* __restrict__ sfa,+ const unsigned char* __restrict__ sfb,+ __half* __restrict__ c,+ int M, int K, int L, int B_rows,+ int sfa_rest_m, int sfa_rest_k, int sfb_rest_m, int sfb_rest_k,+ int64_t a_s0, int64_t a_s1, int64_t a_s2,+ int64_t b_s0, int64_t b_s1, int64_t b_s2,+ int64_t sfa_s0, int64_t sfa_s1, int64_t sfa_s2, int64_t sfa_s3, int64_t sfa_s4, int64_t sfa_s5,+ int64_t sfb_s0, int64_t sfb_s1, int64_t sfb_s2, int64_t sfb_s3, int64_t sfb_s4, int64_t sfb_s5,+ int64_t c_s0, int64_t c_s1, int64_t c_s2+ ) {+ extern __shared__ unsigned char smem[];+ unsigned char* b_shared = smem;+ unsigned char* sfb_shared = smem + (K / 2);- rb = i // 128- ri = i % 128- cb = j // 4- ci = j % 4+ const int warp_id = threadIdx.y;+ const int lane_id = threadIdx.x;+ const int m = blockIdx.x * blockDim.y + warp_id;+ const int l = blockIdx.y;+ const int tid = threadIdx.x + threadIdx.y * blockDim.x;- # Blocking logic from reference- perm = rb * ncb * 128 * 4 + cb * 128 * 4 + ri * 4 + ci+ if (m >= M || l >= L) return;- chunk = perm // 512- in_chunk = perm % 512- d1 = in_chunk // 128- rest = in_chunk % 128- d2 = rest // 4- d3 = rest % 4+ const int K_bytes = K / 2;+ const int K_blocks = K / 16;+ const int b_m = 0;- out_idx = chunk * 512 + d2 * 16 + d1 * 4 + d3+ // Load b (coalesced)+ #pragma unroll 8+ for (int idx = tid; idx < K_bytes; idx += blockDim.x * blockDim.y) {+ b_shared[idx] = __ldg(&b[idx * b_s1 + l * b_s2]);+ }- # Load and store- inp_idx = pid_l * s_il + i * s_im + j * s_ik- out_idx_final = pid_l * s_ol + out_idx * s_oe+ // Precompute scale offsets+ const int mm = m / 128;+ const int mm32 = m % 32;+ const int mm4 = (m % 128) / 32;+ const int64_t sfa_m_part = mm32 * sfa_s0 + mm4 * sfa_s1 + mm * sfa_s2 + l * sfa_s5;- val = tl.load(inp + inp_idx, mask=mask)- tl.store(out + out_idx_final, val, mask=mask)-- def transform_scales_gpu(tensor):- """GPU-based scale transformation"""- M, K, L = tensor.shape- mk = M * K- result = torch.empty((L, mk), dtype=tensor.dtype, device='cuda')+ const int64_t sfb_l_offset = l * sfb_s5;+ #pragma unroll 8+ for (int idx = tid; idx < K_blocks; idx += blockDim.x * blockDim.y) {+ const int kk = idx / 4;+ const int kk4 = idx % 4;+ const int64_t sfb_offset = b_m * sfb_s0 + kk4 * sfb_s3 + kk * sfb_s4 + sfb_l_offset;+ sfb_shared[idx] = __ldg(&sfb[sfb_offset]);+ }+ __syncthreads();- t = tensor.cuda() if not tensor.is_cuda else tensor-- BLK = 256- grid = (L, (mk + BLK - 1) // BLK)-- blocked_transform_kernel[grid](- t, result, M, K, L,- t.stride(0), t.stride(1), t.stride(2),- result.stride(0), result.stride(1),- BLK=BLK- )-- return [result[i] for i in range(L)]-- # =============================================================================- # CUDA GRAPH OPTIMIZATION- # =============================================================================-- _graph_cache = {}-- class CUDAGraphExecutor:- """Captures CUDA graph to eliminate kernel launch overhead"""- def __init__(self, M, K, L):- self.M = M- self.K = K- self.L = L- self.graph = None- self.static_a = None- self.static_b = None- self.static_sfa_list = None- self.static_sfb_list = None- self.static_c = None+ // === VECTORIZED: Process 8 bytes per iteration ===+ float thread_acc = 0.0f;+ #pragma unroll 1+ for (int k_byte = lane_id * 8; k_byte < K_bytes; k_byte += warpSize * 8) {+ const int64_t a_offset = m * a_s0 + k_byte * a_s1 + l * a_s2;- def capture(self, a, b, sfa_list, sfb_list, c):- """Capture CUDA graph"""- # Warmup- for _ in range(3):- for i in range(self.L):- res = torch._scaled_mm(- a[:, :, i],- b[:, :, i].transpose(0, 1),- sfa_list[i],- sfb_list[i],- bias=None,- out_dtype=torch.float16,- )- c[:, 0, i] = res[:, 0]- torch.cuda.synchronize()+ // Load 8 bytes (64 bits) at once+ uint2 a_vec = make_uint2(0, 0);+ if (k_byte + 8 <= K_bytes) {+ a_vec = *reinterpret_cast<const uint2*>(&a[a_offset]);+ } else {+ a_vec.x = (k_byte < K_bytes) ? *reinterpret_cast<const uint*>(&a[a_offset]) : 0;+ a_vec.y = (k_byte + 4 < K_bytes) ? *reinterpret_cast<const uint*>(&a[a_offset + 4]) : 0;+ }- # Create static tensors- self.static_a = a.clone()- self.static_b = b.clone()- self.static_sfa_list = [s.clone() for s in sfa_list]- self.static_sfb_list = [s.clone() for s in sfb_list]- self.static_c = c.clone()-- # Capture- self.graph = torch.cuda.CUDAGraph()- with torch.cuda.graph(self.graph):- for i in range(self.L):- res = torch._scaled_mm(- self.static_a[:, :, i],- self.static_b[:, :, i].transpose(0, 1),- self.static_sfa_list[i],- self.static_sfb_list[i],- bias=None,- out_dtype=torch.float16,- )- self.static_c[:, 0, i] = res[:, 0]-- return self+ // Process 16 FP4 values+ #pragma unroll 8+ for (int i = 0; i < 8; i++) {+ const int byte_idx = k_byte + i;+ if (byte_idx >= K_bytes) break;++ const int k = byte_idx * 2;+ const int k_block = byte_idx >> 3;++ const unsigned char a_val = reinterpret_cast<unsigned char*>(&a_vec)[i];+ const unsigned char b_val = b_shared[byte_idx];++ const float a_fp4_0 = dequant_fp4_e2m1(a_val, 0);+ const float a_fp4_1 = (k + 1 < K) ? dequant_fp4_e2m1(a_val, 1) : 0.0f;+ const float b_fp4_0 = dequant_fp4_e2m1(b_val, 0);+ const float b_fp4_1 = dequant_fp4_e2m1(b_val, 1);++ const int kk = k_block / 4;+ const int kk4 = k_block % 4;+ const int64_t sfa_offset = sfa_m_part + kk4 * sfa_s3 + kk * sfa_s4;++ // Bounds check (computational equivalence)+ if (sfa_offset >= 0) {+ const float scale_a = dequant_fp8_e4m3(__ldg(&sfa[sfa_offset]));+ const float scale_b = dequant_fp8_e4m3(sfb_shared[k_block]);++ const float a_scaled_0 = a_fp4_0 * scale_a;+ const float a_scaled_1 = a_fp4_1 * scale_a;+ const float b_scaled_0 = b_fp4_0 * scale_b;+ const float b_scaled_1 = b_fp4_1 * scale_b;++ thread_acc += a_scaled_0 * b_scaled_0;+ if (k + 1 < K) thread_acc += a_scaled_1 * b_scaled_1;+ }+ }+ }- def execute(self, a, b, sfa_list, sfb_list, c):- """Execute with graph"""- self.static_a.copy_(a, non_blocking=True)- self.static_b.copy_(b, non_blocking=True)- for i in range(self.L):- self.static_sfa_list[i].copy_(sfa_list[i], non_blocking=True)- self.static_sfb_list[i].copy_(sfb_list[i], non_blocking=True)-- self.graph.replay()-- c.copy_(self.static_c, non_blocking=True)- return c+ // Warp reduction+ #pragma unroll+ for (int offset = 16; offset > 0; offset >>= 1) {+ thread_acc += __shfl_down_sync(0xFFFFFFFF, thread_acc, offset);+ }++ if (lane_id == 0) {+ const int64_t c_offset = m * c_s0 + l * c_s2;+ c[c_offset] = __float2half(thread_acc);+ }+ }- # =============================================================================- # MAIN KERNEL- # =============================================================================+ // === PyTorch wrapper (matches original signature) ===- def custom_kernel(data: input_t) -> output_t:- """- Optimized NVFP4 batched GEMV kernel.+ torch::Tensor nvfp4_gemv(+ torch::Tensor a,+ torch::Tensor b,+ torch::Tensor sfa_permuted,+ torch::Tensor sfb_permuted,+ torch::Tensor c+ ) {+ TORCH_CHECK(a.device().is_cuda(), "tensors must be CUDA");+ TORCH_CHECK(sfa_permuted.dim() == 6, "sfa_permuted must be 6D");+ TORCH_CHECK(sfb_permuted.dim() == 6, "sfb_permuted must be 6D");- KEY OPTIMIZATIONS:- 1. GPU-based scale transformation (parallel across L batches)- 2. CUDA graphs (eliminates kernel launch overhead)- 3. Optimized memory access patterns+ unsigned char* a_ptr = reinterpret_cast<unsigned char*>(a.data_ptr());+ unsigned char* b_ptr = reinterpret_cast<unsigned char*>(b.data_ptr());+ unsigned char* sfa_ptr = reinterpret_cast<unsigned char*>(sfa_permuted.data_ptr());+ unsigned char* sfb_ptr = reinterpret_cast<unsigned char*>(sfb_permuted.data_ptr());- LIMITATIONS:- - Still uses torch._scaled_mm which doesn't use Blackwell tensor cores- - To go faster: Need CUTLASS with tcgen05 instructions (multi-file setup)+ int M = a.size(0);+ int K_bytes = a.size(1);+ int L = a.size(2);+ int K = K_bytes * 2;+ int B_rows = b.size(0);- Expected speedup: 1.5-2.5x over reference- """- a, b, sfa, sfb, sfa_permuted, sfb_permuted, c = data+ int sfa_dim2 = sfa_permuted.size(2);+ int sfa_dim4 = sfa_permuted.size(4);+ int sfb_dim2 = sfb_permuted.size(2);+ int sfb_dim4 = sfb_permuted.size(4);- M = a.size(0)- K = a.size(1)- L = a.size(2)+ dim3 block(32, 16);+ dim3 grid((M + 15) / 16, L);+ size_t smem_size = K_bytes + K / 16;- # Transform scales using the CORRECT format for torch._scaled_mm- # (Not the permuted format - that's for CUTLASS)- sfa_list = transform_scales_gpu(sfa)- sfb_list = transform_scales_gpu(sfb)+ nvfp4_gemv_vectorized<<<grid, block, smem_size>>>(+ a_ptr, b_ptr, sfa_ptr, sfb_ptr,+ reinterpret_cast<__half*>(c.data_ptr<at::Half>()),+ M, K, L, B_rows,+ sfa_dim2, sfa_dim4, sfb_dim2, sfb_dim4,+ a.stride(0), a.stride(1), a.stride(2),+ b.stride(0), b.stride(1), b.stride(2),+ sfa_permuted.stride(0), sfa_permuted.stride(1), sfa_permuted.stride(2),+ sfa_permuted.stride(3), sfa_permuted.stride(4), sfa_permuted.stride(5),+ sfb_permuted.stride(0), sfb_permuted.stride(1), sfb_permuted.stride(2),+ sfb_permuted.stride(3), sfb_permuted.stride(4), sfb_permuted.stride(5),+ c.stride(0), c.stride(1), c.stride(2)+ );- # Use CUDA graphs for repeated calls- cache_key = (M, K, L)+ cudaError_t err = cudaGetLastError();+ TORCH_CHECK(err == cudaSuccess, "CUDA error: ", cudaGetErrorString(err));- if cache_key not in _graph_cache:- # First time - capture graph- executor = CUDAGraphExecutor(M, K, L)- executor.capture(a, b, sfa_list, sfb_list, c)- _graph_cache[cache_key] = executor- else:- # Reuse cached graph- executor = _graph_cache[cache_key]-- result = executor.execute(a, b, sfa_list, sfb_list, c)-- return resultNo newline at end of file+ return c;+ }+ """++ nvfp4_gemv_cpp = """+ #include <torch/extension.h>+ torch::Tensor nvfp4_gemv(+ torch::Tensor a,+ torch::Tensor b,+ torch::Tensor sfa,+ torch::Tensor sfb,+ torch::Tensor c+ );+ """++ nvfp4_module = load_inline(+ name='nvfp4_gemv',+ cpp_sources=nvfp4_gemv_cpp,+ cuda_sources=nvfp4_gemv_cuda,+ functions=['nvfp4_gemv'],+ verbose=False,+ extra_cuda_cflags=['-O3', '--use_fast_math', '-arch=sm_100']+ )++ def custom_kernel(data: input_t) -> output_t: # type: ignore+ """+ Vectorized NVFP4 GEMV with shared memory - 8x fewer K-loop iterations+ """+ a, b, _, _, sfa_permuted, sfb_permuted, c = data+ return nvfp4_module.nvfp4_gemv(a, b, sfa_permuted, sfb_permuted, c)No newline at end of file
scrolls · 419 diff lines total
Best evidence level for this revision: reported
JSON