submission 108659
currybab · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 550 lines, June 9 Researcher Reciprocity License v1.0.
submission_l1_specialized.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-108659?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:0ac3f8a6a8cdfc94ebe8540e19daec3852625765dc2375653d369ed2fe232f6b
license declaredunknown
license concludedunknown
authorscurrybab
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
fp8
__device__ __forceinline__ __half decode_fp8_h(__nv_fp8_storage_t v) {shared-memory
__shared__ float smem_acc[4];vector-width = uint4
uint4 a_vec0, a_vec1, b_vec0, b_vec1;Kernel source
submission_l1_specialized.py550 lines
import torch
import sys
import os
import random
import re
import numpy as np
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
cuda_kernel_code = """
#include <cuda_fp4.h>
#include <cuda_fp8.h>
#include <cuda_fp16.h>
#define K_THREADS 32
#define K_L1_THREADS 128
#define WARP_SIZE 32
// Shared memory B caching parameters
#define M_TILE 8 // Number of M rows per block
#define SMEM_K_THREADS 32 // K-direction threads per M row
__device__ __forceinline__ __half decode_fp8_h(__nv_fp8_storage_t v) {
return __half(__nv_cvt_fp8_to_halfraw(v, __NV_E4M3));
}
__device__ __forceinline__ __half2 decode_fp4_h2(__nv_fp4x2_storage_t v) {
return __half2(__nv_cvt_fp4x2_to_halfraw2(v, __NV_E2M1));
}
// =================================================================================================
// Kernel: Specialized L1 Kernel (Optimized for L=1, K=16384)
// =================================================================================================
// Based on v2_merge's gemv_kernel_l1, but with 'l' loop and offsets removed/simplified.
// Hardcoded for L=1.
__global__ void gemv_kernel_l1_specialized(
const __nv_fp4x2_storage_t* __restrict__ A,
const __nv_fp4x2_storage_t* __restrict__ B,
const __nv_fp8_storage_t* __restrict__ sfa,
const __nv_fp8_storage_t* __restrict__ sfb,
half* __restrict__ C,
int M, int K // L is implicitly 1
) {
int m = blockIdx.x;
// L is always 0, so l_offset is 0.
int tid_k = threadIdx.x;
// Shared memory for inter-warp reduction
// K_L1_THREADS / WARP_SIZE = 128 / 32 = 4
__shared__ float smem_acc[4];
if (m >= M) return;
const int K_sf = K / 16;
const int K_half = K / 2;
// Simplified offsets (L=0)
const size_t seg_a = (size_t)m * K_half;
// seg_b is 0
const size_t seg_sfa = (size_t)m * K_sf;
// seg_sfb is 0
float acc = 0.0f;
// 4 blocks per thread iteration (Unroll factor 4)
int k_base = tid_k * 4;
// Pre-load variables
uint4 a_vec0, a_vec1, b_vec0, b_vec1;
uint sfa_packed, sfb_packed;
// Initial Load
if (k_base < K_sf) {
sfa_packed = *reinterpret_cast<const uint*>(&sfa[seg_sfa + k_base]);
sfb_packed = *reinterpret_cast<const uint*>(&sfb[k_base]); // sfb is at offset 0
a_vec0 = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8]);
a_vec1 = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8 + 16]); // +16 uint4 offset? No.
// Wait, previous code: a_vec1 = ... &A[seg_a + k_base * 8 + 16]
// A is __nv_fp4x2_storage_t* (1 byte).
// k_base * 8 -> 8 bytes per block.
// a_vec0 loads 16 bytes (2 blocks: k_base, k_base+1).
// a_vec1 loads 16 bytes (2 blocks: k_base+2, k_base+3).
// 2 blocks = 16 bytes.
// So offset for a_vec1 should be + 16 bytes from a_vec0 address.
// Pointer arithmetic: A + ...
// &A[...] returns address.
// k_base*8 is byte offset? No, A is typed pointer.
// sizeof(__nv_fp4x2_storage_t) = 1.
// So A[...] is byte addressing basically.
// k_base blocks * 8 bytes/block = k_base*8 bytes.
// a_vec0 reads bytes [k_base*8 ... k_base*8+15].
// a_vec1 reads bytes [k_base*8+16 ... k_base*8+31].
// So +16 is correct.
a_vec1 = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8 + 16]);
b_vec0 = *reinterpret_cast<const uint4*>(&B[k_base * 8]);
b_vec1 = *reinterpret_cast<const uint4*>(&B[k_base * 8 + 16]);
}
// Main Loop
// Stride: K_L1_THREADS * 4 = 128 * 4 = 512 blocks
int stride_blocks = K_L1_THREADS * 4;
for (; k_base < K_sf; k_base += stride_blocks) {
uint4 curr_a0 = a_vec0, curr_a1 = a_vec1;
uint4 curr_b0 = b_vec0, curr_b1 = b_vec1;
uint curr_sfa = sfa_packed, curr_sfb = sfb_packed;
// Prefetch next iteration
int next_k = k_base + stride_blocks;
if (next_k < K_sf) {
sfa_packed = *reinterpret_cast<const uint*>(&sfa[seg_sfa + next_k]);
sfb_packed = *reinterpret_cast<const uint*>(&sfb[next_k]);
a_vec0 = *reinterpret_cast<const uint4*>(&A[seg_a + next_k * 8]);
a_vec1 = *reinterpret_cast<const uint4*>(&A[seg_a + next_k * 8 + 16]);
b_vec0 = *reinterpret_cast<const uint4*>(&B[next_k * 8]);
b_vec1 = *reinterpret_cast<const uint4*>(&B[next_k * 8 + 16]);
}
const __nv_fp8_storage_t* sfa_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&curr_sfa);
const __nv_fp8_storage_t* sfb_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&curr_sfb);
// Process first 2 blocks (k_base, k_base+1)
const __nv_fp4x2_storage_t* a_fp4x2_0 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_a0);
const __nv_fp4x2_storage_t* b_fp4x2_0 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_b0);
#pragma unroll
for (int i = 0; i < 2; i++) {
if (k_base + i >= K_sf) break; // Safety check, usually not needed if K divisible
float scale_ab = __half2float(decode_fp8_h(sfa_bytes[i])) * __half2float(decode_fp8_h(sfb_bytes[i]));
__half2 block_acc_h2 = __float2half2_rn(0.0f);
#pragma unroll
for (int j = 0; j < 8; j++) {
__half2 a2 = decode_fp4_h2(a_fp4x2_0[i * 8 + j]);
__half2 b2 = decode_fp4_h2(b_fp4x2_0[i * 8 + j]);
block_acc_h2 = __hfma2(a2, b2, block_acc_h2);
}
float2 pf = __half22float2(block_acc_h2);
acc = __fmaf_rn(pf.x + pf.y, scale_ab, acc);
}
// Process next 2 blocks (k_base+2, k_base+3)
const __nv_fp4x2_storage_t* a_fp4x2_1 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_a1);
const __nv_fp4x2_storage_t* b_fp4x2_1 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_b1);
#pragma unroll
for (int i = 0; i < 2; i++) {
if (k_base + 2 + i >= K_sf) break;
float scale_ab = __half2float(decode_fp8_h(sfa_bytes[2 + i])) * __half2float(decode_fp8_h(sfb_bytes[2 + i]));
__half2 block_acc_h2 = __float2half2_rn(0.0f);
#pragma unroll
for (int j = 0; j < 8; j++) {
__half2 a2 = decode_fp4_h2(a_fp4x2_1[i * 8 + j]);
__half2 b2 = decode_fp4_h2(b_fp4x2_1[i * 8 + j]);
block_acc_h2 = __hfma2(a2, b2, block_acc_h2);
}
float2 pf = __half22float2(block_acc_h2);
acc = __fmaf_rn(pf.x + pf.y, scale_ab, acc);
}
}
// Warp reduction
#pragma unroll
for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
acc += __shfl_down_sync(0xffffffff, acc, offset);
}
int warp_id = tid_k / WARP_SIZE; // 0..3
int lane_id = tid_k % WARP_SIZE;
if (lane_id == 0) {
smem_acc[warp_id] = acc;
}
__syncthreads();
if (warp_id == 0) {
float val = (tid_k < 4) ? smem_acc[tid_k] : 0.0f;
val += __shfl_down_sync(0xffffffff, val, 2);
val += __shfl_down_sync(0xffffffff, val, 1);
if (tid_k == 0) {
size_t c_idx = (size_t)m; // L=0
C[c_idx] = __float2half(val);
}
}
}
// Shared memory B caching kernel (For Large L)
// Copy from v2_merge
__global__ void gemv_kernel_smem(
const __nv_fp4x2_storage_t* __restrict__ A,
const __nv_fp4x2_storage_t* __restrict__ B,
const __nv_fp8_storage_t* __restrict__ sfa,
const __nv_fp8_storage_t* __restrict__ sfb,
half* __restrict__ C,
int M, int K, int L
) {
int m_tile_idx = blockIdx.x;
int l = blockIdx.y;
int tid_k = threadIdx.x; // K-direction
int tid_m = threadIdx.y; // M-direction within tile
int m = m_tile_idx * M_TILE + tid_m;
if (l >= L) return;
const int K_sf = K / 16;
const int K_half = K / 2;
extern __shared__ char shared_mem[];
__nv_fp4x2_storage_t* s_B = (__nv_fp4x2_storage_t*)shared_mem;
__nv_fp8_storage_t* s_sfb = (__nv_fp8_storage_t*)(s_B + K_half);
const size_t seg_b = (size_t)128 * K_half * l;
const size_t seg_sfb = (size_t)128 * K_sf * l;
int total_threads = blockDim.x * blockDim.y;
int linear_tid = tid_m * blockDim.x + tid_k;
for (int i = linear_tid; i < K_half / 16; i += total_threads) {
uint4 b_vec = *reinterpret_cast<const uint4*>(&B[seg_b + i * 16]);
*reinterpret_cast<uint4*>(&s_B[i * 16]) = b_vec;
}
for (int i = (K_half / 16) * 16 + linear_tid; i < K_half; i += total_threads) {
s_B[i] = B[seg_b + i];
}
for (int i = linear_tid; i < K_sf / 16; i += total_threads) {
uint4 sf_vec = *reinterpret_cast<const uint4*>(&sfb[seg_sfb + i * 16]);
*reinterpret_cast<uint4*>(&s_sfb[i * 16]) = sf_vec;
}
for (int i = (K_sf / 16) * 16 + linear_tid; i < K_sf; i += total_threads) {
s_sfb[i] = sfb[seg_sfb + i];
}
__syncthreads();
if (m >= M) return;
const size_t seg_a = (size_t)M * K_half * l + (size_t)m * K_half;
const size_t seg_sfa = (size_t)M * K_sf * l + (size_t)m * K_sf;
float acc = 0.0f;
for (int k_base = tid_k * 2; k_base < K_sf; k_base += SMEM_K_THREADS * 2) {
__nv_fp8x2_storage_t sfa_pair = *reinterpret_cast<const __nv_fp8x2_storage_t*>(&sfa[seg_sfa + k_base]);
__nv_fp8x2_storage_t sfb_pair = *reinterpret_cast<const __nv_fp8x2_storage_t*>(&s_sfb[k_base]);
const __nv_fp8_storage_t* sfa_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&sfa_pair);
const __nv_fp8_storage_t* sfb_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&sfb_pair);
uint4 a_vec = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8]);
uint4 b_vec = *reinterpret_cast<const uint4*>(&s_B[k_base * 8]);
const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_vec);
const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_vec);
#pragma unroll
for (int i = 0; i < 2; i++) {
int k = k_base + i;
if (k >= K_sf) break;
float scale_a = __half2float(decode_fp8_h(sfa_bytes[i]));
float scale_b = __half2float(decode_fp8_h(sfb_bytes[i]));
float scale_ab = scale_a * scale_b;
float block_sum = 0.0f;
#pragma unroll
for (int j = 0; j < 8; j++) {
__half2 a2 = decode_fp4_h2(a_fp4x2[i * 8 + j]);
__half2 b2 = decode_fp4_h2(b_fp4x2[i * 8 + j]);
__half2 prod = __hmul2(a2, b2);
float2 pf = __half22float2(prod);
block_sum += (pf.x + pf.y);
}
acc = __fmaf_rn(block_sum, scale_ab, acc);
}
}
#pragma unroll
for (int offset = WARP_SIZE >> 1; offset > 0; offset >>= 1) {
acc += __shfl_down_sync(0xffffffff, acc, offset);
}
if (tid_k == 0) {
size_t c_idx = (size_t)m + (size_t)l * (size_t)M;
C[c_idx] = __float2half(acc);
}
}
// Original L1 kernel (for L=2,3)
__global__ void gemv_kernel_l1(
const __nv_fp4x2_storage_t* __restrict__ A,
const __nv_fp4x2_storage_t* __restrict__ B,
const __nv_fp8_storage_t* __restrict__ sfa,
const __nv_fp8_storage_t* __restrict__ sfb,
half* __restrict__ C,
int M, int K, int L
) {
int m = blockIdx.x;
int l = blockIdx.y;
int tid_k = threadIdx.x;
__shared__ float smem_acc[4];
if (m >= M || l >= L) return;
const int K_sf = K / 16;
const int K_half = K / 2;
const size_t seg_a = (size_t)M * K_half * l + (size_t)m * K_half;
const size_t seg_b = (size_t)128 * K_half * l;
const size_t seg_sfa = (size_t)M * K_sf * l + (size_t)m * K_sf;
const size_t seg_sfb = (size_t)128 * K_sf * l;
float acc = 0.0f;
int k_base = tid_k * 4;
uint4 a_vec0, a_vec1, b_vec0, b_vec1;
uint sfa_packed, sfb_packed;
if (k_base < K_sf) {
sfa_packed = *reinterpret_cast<const uint*>(&sfa[seg_sfa + k_base]);
sfb_packed = *reinterpret_cast<const uint*>(&sfb[seg_sfb + k_base]);
a_vec0 = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8]);
a_vec1 = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8 + 16]);
b_vec0 = *reinterpret_cast<const uint4*>(&B[seg_b + k_base * 8]);
b_vec1 = *reinterpret_cast<const uint4*>(&B[seg_b + k_base * 8 + 16]);
}
for (; k_base < K_sf; k_base += K_L1_THREADS * 4) {
uint4 curr_a0 = a_vec0, curr_a1 = a_vec1;
uint4 curr_b0 = b_vec0, curr_b1 = b_vec1;
uint curr_sfa = sfa_packed, curr_sfb = sfb_packed;
int next_k = k_base + K_L1_THREADS * 4;
if (next_k < K_sf) {
sfa_packed = *reinterpret_cast<const uint*>(&sfa[seg_sfa + next_k]);
sfb_packed = *reinterpret_cast<const uint*>(&sfb[seg_sfb + next_k]);
a_vec0 = *reinterpret_cast<const uint4*>(&A[seg_a + next_k * 8]);
a_vec1 = *reinterpret_cast<const uint4*>(&A[seg_a + next_k * 8 + 16]);
b_vec0 = *reinterpret_cast<const uint4*>(&B[seg_b + next_k * 8]);
b_vec1 = *reinterpret_cast<const uint4*>(&B[seg_b + next_k * 8 + 16]);
}
const __nv_fp8_storage_t* sfa_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&curr_sfa);
const __nv_fp8_storage_t* sfb_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&curr_sfb);
const __nv_fp4x2_storage_t* a_fp4x2_0 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_a0);
const __nv_fp4x2_storage_t* b_fp4x2_0 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_b0);
#pragma unroll
for (int i = 0; i < 2; i++) {
if (k_base + i >= K_sf) break;
float scale_ab = __half2float(decode_fp8_h(sfa_bytes[i])) * __half2float(decode_fp8_h(sfb_bytes[i]));
__half2 block_acc_h2 = __float2half2_rn(0.0f);
#pragma unroll
for (int j = 0; j < 8; j++) {
__half2 a2 = decode_fp4_h2(a_fp4x2_0[i * 8 + j]);
__half2 b2 = decode_fp4_h2(b_fp4x2_0[i * 8 + j]);
block_acc_h2 = __hfma2(a2, b2, block_acc_h2);
}
float2 pf = __half22float2(block_acc_h2);
acc = __fmaf_rn(pf.x + pf.y, scale_ab, acc);
}
const __nv_fp4x2_storage_t* a_fp4x2_1 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_a1);
const __nv_fp4x2_storage_t* b_fp4x2_1 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&curr_b1);
#pragma unroll
for (int i = 0; i < 2; i++) {
if (k_base + 2 + i >= K_sf) break;
float scale_ab = __half2float(decode_fp8_h(sfa_bytes[2 + i])) * __half2float(decode_fp8_h(sfb_bytes[2 + i]));
__half2 block_acc_h2 = __float2half2_rn(0.0f);
#pragma unroll
for (int j = 0; j < 8; j++) {
__half2 a2 = decode_fp4_h2(a_fp4x2_1[i * 8 + j]);
__half2 b2 = decode_fp4_h2(b_fp4x2_1[i * 8 + j]);
block_acc_h2 = __hfma2(a2, b2, block_acc_h2);
}
float2 pf = __half22float2(block_acc_h2);
acc = __fmaf_rn(pf.x + pf.y, scale_ab, acc);
}
}
#pragma unroll
for (int offset = WARP_SIZE / 2; offset > 0; offset /= 2) {
acc += __shfl_down_sync(0xffffffff, acc, offset);
}
int warp_id = tid_k / WARP_SIZE;
int lane_id = tid_k % WARP_SIZE;
if (lane_id == 0) smem_acc[warp_id] = acc;
__syncthreads();
if (warp_id == 0) {
float val = (tid_k < 4) ? smem_acc[tid_k] : 0.0f;
val += __shfl_down_sync(0xffffffff, val, 2);
val += __shfl_down_sync(0xffffffff, val, 1);
if (tid_k == 0) {
size_t c_idx = (size_t)m + (size_t)l * (size_t)M;
C[c_idx] = __float2half(val);
}
}
}
__global__ void gemv_kernel(
const __nv_fp4x2_storage_t* __restrict__ A,
const __nv_fp4x2_storage_t* __restrict__ B,
const __nv_fp8_storage_t* __restrict__ sfa,
const __nv_fp8_storage_t* __restrict__ sfb,
half* __restrict__ C,
int M, int K, int L
) {
int m = blockIdx.x;
int l_base = blockIdx.y * blockDim.y;
int tid_k = threadIdx.x; // K 방향
int tid_l = threadIdx.y; // L 방향
int l = l_base + tid_l;
if (m >= M || l >= L) return;
const int K_sf = K / 16;
const int K_half = K / 2;
const size_t seg_a = (size_t)M * K_half * l + (size_t)m * K_half;
const size_t seg_b = (size_t)128 * K_half * l;
const size_t seg_sfa = (size_t)M * K_sf * l + (size_t)m * K_sf;
const size_t seg_sfb = (size_t)128 * K_sf * l;
float acc = 0.0f;
// 2개 k 블록씩 처리 (원래 로직)
for (int k_base = tid_k * 2; k_base < K_sf; k_base += K_THREADS * 2) {
__nv_fp8x2_storage_t sfa_pair = *reinterpret_cast<const __nv_fp8x2_storage_t*>(&sfa[seg_sfa + k_base]);
__nv_fp8x2_storage_t sfb_pair = *reinterpret_cast<const __nv_fp8x2_storage_t*>(&sfb[seg_sfb + k_base]);
const __nv_fp8_storage_t* sfa_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&sfa_pair);
const __nv_fp8_storage_t* sfb_bytes = reinterpret_cast<const __nv_fp8_storage_t*>(&sfb_pair);
uint4 a_vec = *reinterpret_cast<const uint4*>(&A[seg_a + k_base * 8]);
uint4 b_vec = *reinterpret_cast<const uint4*>(&B[seg_b + k_base * 8]);
const __nv_fp4x2_storage_t* a_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&a_vec);
const __nv_fp4x2_storage_t* b_fp4x2 = reinterpret_cast<const __nv_fp4x2_storage_t*>(&b_vec);
#pragma unroll
for (int i = 0; i < 2; i++) {
int k = k_base + i;
if (k >= K_sf) break;
float scale_a = __half2float(decode_fp8_h(sfa_bytes[i]));
float scale_b = __half2float(decode_fp8_h(sfb_bytes[i]));
float scale_ab = scale_a * scale_b;
float block_sum = 0.0f;
#pragma unroll
for (int j = 0; j < 8; j++) {
__half2 a2 = decode_fp4_h2(a_fp4x2[i * 8 + j]);
__half2 b2 = decode_fp4_h2(b_fp4x2[i * 8 + j]);
__half2 prod = __hmul2(a2, b2);
float2 pf = __half22float2(prod);
block_sum += (pf.x + pf.y);
}
acc = __fmaf_rn(block_sum, scale_ab, acc);
}
}
// Warp reduction
#pragma unroll
for (int offset = WARP_SIZE >> 1; offset > 0; offset >>= 1) {
acc += __shfl_down_sync(0xffffffff, acc, offset);
}
size_t c_idx = (size_t)m + (size_t)l * (size_t)M;
C[c_idx] = __float2half(acc);
}
torch::Tensor gemv_cuda(torch::Tensor A, torch::Tensor B, torch::Tensor C, torch::Tensor sfa, torch::Tensor sfb) {
int M = A.size(0);
int K = A.size(1) * 2;
int L = A.size(2);
if (L == 1) {
// Specialized L1 kernel (No L overhead)
dim3 threadsPerBlock(K_L1_THREADS);
dim3 blocksPerGrid(M, 1); // L=1
gemv_kernel_l1_specialized<<<blocksPerGrid, threadsPerBlock>>>(
reinterpret_cast<__nv_fp4x2_storage_t*>(A.data_ptr()),
reinterpret_cast<__nv_fp4x2_storage_t*>(B.data_ptr()),
reinterpret_cast<__nv_fp8_storage_t*>(sfa.data_ptr()),
reinterpret_cast<__nv_fp8_storage_t*>(sfb.data_ptr()),
reinterpret_cast<half*>(C.data_ptr()),
M,
K
);
} else if (L < 4) {
// Use L1 kernel with 4-block pipelining for small L
dim3 threadsPerBlock(K_L1_THREADS);
dim3 blocksPerGrid(M, L);
gemv_kernel_l1<<<blocksPerGrid, threadsPerBlock>>>(
reinterpret_cast<__nv_fp4x2_storage_t*>(A.data_ptr()),
reinterpret_cast<__nv_fp4x2_storage_t*>(B.data_ptr()),
reinterpret_cast<__nv_fp8_storage_t*>(sfa.data_ptr()),
reinterpret_cast<__nv_fp8_storage_t*>(sfb.data_ptr()),
reinterpret_cast<half*>(C.data_ptr()),
M,
K,
L
);
} else if (L >= 8) {
// Use shared memory B caching for large L (better B reuse)
int K_half = K / 2;
int K_sf = K / 16;
size_t shared_mem_size = K_half * sizeof(unsigned char) + K_sf * sizeof(unsigned char);
dim3 threadsPerBlock(SMEM_K_THREADS, M_TILE);
dim3 blocksPerGrid((M + M_TILE - 1) / M_TILE, L);
gemv_kernel_smem<<<blocksPerGrid, threadsPerBlock, shared_mem_size>>>(
reinterpret_cast<__nv_fp4x2_storage_t*>(A.data_ptr()),
reinterpret_cast<__nv_fp4x2_storage_t*>(B.data_ptr()),
reinterpret_cast<__nv_fp8_storage_t*>(sfa.data_ptr()),
reinterpret_cast<__nv_fp8_storage_t*>(sfb.data_ptr()),
reinterpret_cast<half*>(C.data_ptr()),
M,
K,
L
);
} else {
// L = 4~7: use original kernel
int L_TILE_SIZE = 4;
dim3 threadsPerBlock(K_THREADS, L_TILE_SIZE);
dim3 blocksPerGrid(M, (L + L_TILE_SIZE - 1) / L_TILE_SIZE);
gemv_kernel<<<blocksPerGrid, threadsPerBlock>>>(
reinterpret_cast<__nv_fp4x2_storage_t*>(A.data_ptr()),
reinterpret_cast<__nv_fp4x2_storage_t*>(B.data_ptr()),
reinterpret_cast<__nv_fp8_storage_t*>(sfa.data_ptr()),
reinterpret_cast<__nv_fp8_storage_t*>(sfb.data_ptr()),
reinterpret_cast<half*>(C.data_ptr()),
M,
K,
L
);
}
return C;
}
"""
cpp_code = """
#include <torch/extension.h>
torch::Tensor gemv_cuda(torch::Tensor A, torch::Tensor B, torch::Tensor C, torch::Tensor sfa, torch::Tensor sfb);
"""
gemv_module = load_inline(
name='gemv_cuda_l1_spec',
cpp_sources=cpp_code,
cuda_sources=cuda_kernel_code,
functions=['gemv_cuda'],
with_cuda=True,
extra_cuda_cflags=["-O3", "-use_fast_math", "-gencode=arch=compute_100a,code=sm_100a"],
verbose=False,
)
def gemv(A, B, C, sfa, sfb):
if not A.is_cuda or not B.is_cuda or not C.is_cuda or not sfa.is_cuda or not sfb.is_cuda:
raise RuntimeError("All tensors must be on GPU")
return gemv_module.gemv_cuda(A, B, C, sfa, sfb)
def custom_kernel(
data: input_t,
) -> output_t:
a_ref, b_ref, sfa, sfb, _, _, c_ref = data
sfa = sfa.to(device=torch.cuda.current_device())
sfb = sfb.to(device=torch.cuda.current_device())
return gemv(a_ref, b_ref, c_ref, sfa, sfb)
scrolls · 550 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