submission 104755
mysfi · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 647 lines, June 9 Researcher Reciprocity License v1.0.
submission_inline.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-104755?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:627a9d5c25144bcf8bbb27a3ebe17cd88101be8d4b4009b2fbe902cfb8b49bdd
license declaredunknown
license concludedunknown
authorsmysfi
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
async-copy
cutlass::arch::cp_async<sizeof(uint32_t)>(smem_ptr, gmem_ptr, stage_valid);fp4
using ElementA = cutlass::float_e2m1_t;fp8
using ElementSFA = cutlass::float_e4m3_t;fused-epilogue
void apply_epilogue(shared-memory
__align__(16) ElementA smem_A[kBufferCount][kStageCount][kSmemPerStageA];Kernel source
submission_inline.py647 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
CUTLASS_PATHS = [
"/workspace/cutlass/include",
"/workspace/cutlass/tools/include",
]
# CUDA source code with CUTLASS/CuTe header probes
cuda_src = r"""
#include <torch/extension.h>
#include <c10/cuda/CUDAStream.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include "cutlass/cutlass.h"
#include "cutlass/layout/matrix.h"
#include "cutlass/matrix_coord.h"
#include "cutlass/numeric_types.h"
#include "cutlass/numeric_conversion.h"
#include "cutlass/arch/memory.h"
#include "cutlass/arch/memory_sm80.h"
#include "cutlass/fast_math.h"
#include <cstdint>
namespace nvfp4_gemv_sketch {
using ElementA = cutlass::float_e2m1_t;
using ElementB = cutlass::float_e2m1_t;
using ElementC = cutlass::half_t;
using ElementSFA = cutlass::float_e4m3_t;
using ElementSFB = cutlass::float_e4m3_t;
using ElementAccumulator = float;
static constexpr int kElementsPerAccess = 32;
static constexpr int kSFVecSize = 16;
static constexpr int kSFPerAccess = 2;
static constexpr int kThreadCount = 128;
static constexpr int kThreadsPerRow = 8;
static constexpr int kThreadsPerCol = kThreadCount / kThreadsPerRow;
static constexpr int kStageCount = 4;
static constexpr int kBufferCount = 2;
static constexpr int kPackedElementsA = 2;
static constexpr int kPackedElementsB = 2;
static constexpr int kPackedElements = 2;
static constexpr int kSmemPerStageA = kThreadCount * (kElementsPerAccess / kPackedElementsA);
static constexpr int kSmemPerStageB = kThreadsPerRow * (kElementsPerAccess / kPackedElementsB);
using FragmentA = cutlass::Array<ElementA, kElementsPerAccess>;
using FragmentB = cutlass::Array<ElementB, kElementsPerAccess>;
using FragmentCompute = cutlass::Array<ElementAccumulator, kElementsPerAccess>;
using FragmentSFA = cutlass::Array<ElementSFA, kSFPerAccess>;
using FragmentSFB = cutlass::Array<ElementSFB, kSFPerAccess>;
using FragmentPackedA = cutlass::Array<ElementA, kPackedElements>;
using FragmentPackedB = cutlass::Array<ElementB, kPackedElements>;
static constexpr cutlass::FloatRoundStyle kRound = cutlass::FloatRoundStyle::round_to_nearest;
using A2AccConverter = cutlass::NumericConverter<ElementAccumulator, ElementA, kRound>;
using B2AccConverter = cutlass::NumericConverter<ElementAccumulator, ElementB, kRound>;
using SFA2AccConverter = cutlass::NumericConverter<ElementAccumulator, ElementSFA, kRound>;
using SFB2AccConverter = cutlass::NumericConverter<ElementAccumulator, ElementSFB, kRound>;
struct Params {
int M = 0;
int K = 0;
int batch_count = 0;
ElementA const* ptr_A = nullptr;
int64_t stride_A = 0;
int64_t batch_stride_A = 0;
ElementB const* ptr_B = nullptr;
int64_t batch_stride_B = 0;
ElementC* ptr_C = nullptr;
int64_t batch_stride_C = 0;
ElementSFA const* ptr_SFA = nullptr;
ElementSFB const* ptr_SFB = nullptr;
int64_t batch_stride_SFA = 0;
int64_t batch_stride_SFB = 0;
Params() = default;
};
struct SharedStorage {
__align__(16) ElementA smem_A[kBufferCount][kStageCount][kSmemPerStageA];
__align__(16) ElementB smem_B[kBufferCount][kStageCount][kSmemPerStageB];
__align__(16) ElementSFA smem_SFA[kBufferCount][kStageCount][kThreadCount * kSFPerAccess];
__align__(16) ElementSFB smem_SFB[kBufferCount][kStageCount][kThreadsPerRow * kSFPerAccess];
};
__device__ __forceinline__
int compute_sf_offset_from_k(int global_k) {
int SF_idx = global_k / kSFVecSize;
int block = SF_idx >> 2;
int offset = SF_idx & 0x3;
int result = (block << 9) + offset;
return result;
}
CUTLASS_DEVICE
ElementAccumulator blockscaled_multiply_add(
FragmentA const& fragA,
FragmentB const& fragB,
FragmentSFA const& fragSFA,
FragmentSFB const& fragSFB) {
uint16_t const& src_fragSFA_packed = reinterpret_cast<uint16_t const&>(fragSFA);
uint16_t const& src_fragSFB_packed = reinterpret_cast<uint16_t const&>(fragSFB);
uint32_t const* src_fragA_packed = reinterpret_cast<uint32_t const*>(&fragA);
uint32_t const* src_fragB_packed = reinterpret_cast<uint32_t const*>(&fragB);
__half out_h;
uint16_t* out_fp16 = reinterpret_cast<uint16_t*>(&out_h);
asm volatile( \
"{\n" \
".reg .b8 byte0_0, byte0_1, byte0_2, byte0_3;\n" \
".reg .b8 byte0_4, byte0_5, byte0_6, byte0_7;\n" \
".reg .b8 byte1_0, byte1_1, byte1_2, byte1_3;\n" \
".reg .b8 byte1_4, byte1_5, byte1_6, byte1_7;\n" \
".reg .b8 byte2_0, byte2_1, byte2_2, byte2_3;\n" \
".reg .b8 byte2_4, byte2_5, byte2_6, byte2_7;\n" \
".reg .b8 byte3_0, byte3_1, byte3_2, byte3_3;\n" \
".reg .b8 byte3_4, byte3_5, byte3_6, byte3_7;\n" \
".reg .f16x2 accum_0_0, accum_0_1, accum_0_2, accum_0_3;\n" \
".reg .f16x2 accum_1_0, accum_1_1, accum_1_2, accum_1_3;\n" \
".reg .f16x2 accum_2_0, accum_2_1, accum_2_2, accum_2_3;\n" \
".reg .f16x2 accum_3_0, accum_3_1, accum_3_2, accum_3_3;\n" \
".reg .f16x2 sfa_f16x2;\n" \
".reg .f16x2 sfb_f16x2;\n" \
".reg .f16x2 sf_f16x2;\n" \
".reg .f16x2 cvt_0_0, cvt_0_1, cvt_0_2, cvt_0_3;\n" \
".reg .f16x2 cvt_0_4, cvt_0_5, cvt_0_6, cvt_0_7;\n" \
".reg .f16x2 cvt_1_0, cvt_1_1, cvt_1_2, cvt_1_3;\n" \
".reg .f16x2 cvt_1_4, cvt_1_5, cvt_1_6, cvt_1_7;\n" \
".reg .f16x2 cvt_2_0, cvt_2_1, cvt_2_2, cvt_2_3;\n" \
".reg .f16x2 cvt_2_4, cvt_2_5, cvt_2_6, cvt_2_7;\n" \
".reg .f16x2 cvt_3_0, cvt_3_1, cvt_3_2, cvt_3_3;\n" \
".reg .f16x2 cvt_3_4, cvt_3_5, cvt_3_6, cvt_3_7;\n" \
".reg .f16 result_f16, lane0, lane1;\n" \
".reg .f16x2 mul_f16x2_0, mul_f16x2_1;\n" \
"cvt.rn.f16x2.e4m3x2 sfa_f16x2, %1;\n" \
"cvt.rn.f16x2.e4m3x2 sfb_f16x2, %2;\n" \
"mov.b32 accum_0_0, 0;\n" \
"mov.b32 accum_0_1, 0;\n" \
"mov.b32 accum_0_2, 0;\n" \
"mov.b32 accum_0_3, 0;\n" \
"mov.b32 accum_1_0, 0;\n" \
"mov.b32 accum_1_1, 0;\n" \
"mov.b32 accum_1_2, 0;\n" \
"mov.b32 accum_1_3, 0;\n" \
"mov.b32 accum_2_0, 0;\n" \
"mov.b32 accum_2_1, 0;\n" \
"mov.b32 accum_2_2, 0;\n" \
"mov.b32 accum_2_3, 0;\n" \
"mov.b32 accum_3_0, 0;\n" \
"mov.b32 accum_3_1, 0;\n" \
"mov.b32 accum_3_2, 0;\n" \
"mov.b32 accum_3_3, 0;\n" \
"mul.rn.f16x2 sf_f16x2, sfa_f16x2, sfb_f16x2;\n" \
"mov.b32 {lane0, lane1}, sf_f16x2;\n" \
"mov.b32 mul_f16x2_0, {lane0, lane0};\n" \
"mov.b32 mul_f16x2_1, {lane1, lane1};\n" \
"mov.b32 {byte0_0, byte0_1, byte0_2, byte0_3}, %3;\n" \
"mov.b32 {byte0_4, byte0_5, byte0_6, byte0_7}, %4;\n" \
"mov.b32 {byte1_0, byte1_1, byte1_2, byte1_3}, %5;\n" \
"mov.b32 {byte1_4, byte1_5, byte1_6, byte1_7}, %6;\n" \
"mov.b32 {byte2_0, byte2_1, byte2_2, byte2_3}, %7;\n" \
"mov.b32 {byte2_4, byte2_5, byte2_6, byte2_7}, %8;\n" \
"mov.b32 {byte3_0, byte3_1, byte3_2, byte3_3}, %9;\n" \
"mov.b32 {byte3_4, byte3_5, byte3_6, byte3_7}, %10;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_0_0, byte0_0;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_0_1, byte0_1;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_0_2, byte0_2;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_0_3, byte0_3;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_0_4, byte0_4;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_0_5, byte0_5;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_0_6, byte0_6;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_0_7, byte0_7;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_1_0, byte1_0;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_1_1, byte1_1;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_1_2, byte1_2;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_1_3, byte1_3;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_1_4, byte1_4;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_1_5, byte1_5;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_1_6, byte1_6;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_1_7, byte1_7;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_2_0, byte2_0;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_2_1, byte2_1;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_2_2, byte2_2;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_2_3, byte2_3;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_2_4, byte2_4;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_2_5, byte2_5;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_2_6, byte2_6;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_2_7, byte2_7;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_3_0, byte3_0;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_3_1, byte3_1;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_3_2, byte3_2;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_3_3, byte3_3;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_3_4, byte3_4;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_3_5, byte3_5;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_3_6, byte3_6;\n" \
"cvt.rn.f16x2.e2m1x2 cvt_3_7, byte3_7;\n" \
"fma.rn.f16x2 accum_0_0, cvt_0_0, cvt_0_4, accum_0_0;\n" \
"fma.rn.f16x2 accum_0_1, cvt_0_1, cvt_0_5, accum_0_1;\n" \
"fma.rn.f16x2 accum_0_2, cvt_0_2, cvt_0_6, accum_0_2;\n" \
"fma.rn.f16x2 accum_0_3, cvt_0_3, cvt_0_7, accum_0_3;\n" \
"fma.rn.f16x2 accum_1_0, cvt_1_0, cvt_1_4, accum_1_0;\n" \
"fma.rn.f16x2 accum_1_1, cvt_1_1, cvt_1_5, accum_1_1;\n" \
"fma.rn.f16x2 accum_1_2, cvt_1_2, cvt_1_6, accum_1_2;\n" \
"fma.rn.f16x2 accum_1_3, cvt_1_3, cvt_1_7, accum_1_3;\n" \
"fma.rn.f16x2 accum_2_0, cvt_2_0, cvt_2_4, accum_2_0;\n" \
"fma.rn.f16x2 accum_2_1, cvt_2_1, cvt_2_5, accum_2_1;\n" \
"fma.rn.f16x2 accum_2_2, cvt_2_2, cvt_2_6, accum_2_2;\n" \
"fma.rn.f16x2 accum_2_3, cvt_2_3, cvt_2_7, accum_2_3;\n" \
"fma.rn.f16x2 accum_3_0, cvt_3_0, cvt_3_4, accum_3_0;\n" \
"fma.rn.f16x2 accum_3_1, cvt_3_1, cvt_3_5, accum_3_1;\n" \
"fma.rn.f16x2 accum_3_2, cvt_3_2, cvt_3_6, accum_3_2;\n" \
"fma.rn.f16x2 accum_3_3, cvt_3_3, cvt_3_7, accum_3_3;\n" \
"add.rn.f16x2 accum_0_0, accum_0_0, accum_0_1;\n" \
"add.rn.f16x2 accum_0_2, accum_0_2, accum_0_3;\n" \
"add.rn.f16x2 accum_1_0, accum_1_0, accum_1_1;\n" \
"add.rn.f16x2 accum_1_2, accum_1_2, accum_1_3;\n" \
"add.rn.f16x2 accum_2_0, accum_2_0, accum_2_1;\n" \
"add.rn.f16x2 accum_2_2, accum_2_2, accum_2_3;\n" \
"add.rn.f16x2 accum_3_0, accum_3_0, accum_3_1;\n" \
"add.rn.f16x2 accum_3_2, accum_3_2, accum_3_3;\n" \
"add.rn.f16x2 accum_0_0, accum_0_0, accum_0_2;\n" \
"add.rn.f16x2 accum_1_0, accum_1_0, accum_1_2;\n" \
"add.rn.f16x2 accum_2_0, accum_2_0, accum_2_2;\n" \
"add.rn.f16x2 accum_3_0, accum_3_0, accum_3_2;\n" \
"add.rn.f16x2 accum_0_0, accum_0_0, accum_1_0;\n" \
"add.rn.f16x2 accum_2_0, accum_2_0, accum_3_0;\n" \
"mul.rn.f16x2 accum_0_0, mul_f16x2_0, accum_0_0;\n" \
"mul.rn.f16x2 accum_2_0, mul_f16x2_1, accum_2_0;\n" \
"add.rn.f16x2 accum_0_0, accum_0_0, accum_2_0;\n" \
"mov.b32 {lane0, lane1}, accum_0_0;\n" \
"add.rn.f16 result_f16, lane0, lane1;\n" \
"mov.b16 %0, result_f16;\n" \
"}\n"
: "=h"(out_fp16[0])
: "h"(src_fragSFA_packed), "h"(src_fragSFB_packed),
"r"(src_fragA_packed[0]), "r"(src_fragB_packed[0]),
"r"(src_fragA_packed[1]), "r"(src_fragB_packed[1]),
"r"(src_fragA_packed[2]), "r"(src_fragB_packed[2]),
"r"(src_fragA_packed[3]), "r"(src_fragB_packed[3])
: "memory"
);
return __half2float(out_h);
}
__device__
ElementAccumulator process_tail_elements(
int unrolled_k,
int idx_col_k,
int K,
ElementA const* ptr_A,
ElementB const* ptr_B,
ElementSFA const* ptr_SFA,
ElementSFB const* ptr_SFB)
{
A2AccConverter A_converter;
B2AccConverter B_converter;
SFA2AccConverter SFA_converter;
SFB2AccConverter SFB_converter;
ElementAccumulator accum = ElementAccumulator(0);
for (int k = unrolled_k + idx_col_k * kPackedElementsA;
k < K;
k += kThreadsPerRow * kPackedElementsA) {
int SF_offset_by_k = compute_sf_offset_from_k(k);
ElementSFA sfa = *(ptr_SFA + SF_offset_by_k);
ElementSFB sfb = *(ptr_SFB + SF_offset_by_k);
FragmentPackedA fragA;
FragmentPackedB fragB;
cutlass::arch::global_load<FragmentPackedA, sizeof(FragmentPackedA),
cutlass::arch::CacheOperation::Always>(
fragA,
ptr_A - (idx_col_k * kElementsPerAccess - k) / kPackedElementsA,
true);
cutlass::arch::global_load<FragmentPackedB, sizeof(FragmentPackedB),
cutlass::arch::CacheOperation::Always>(
fragB,
ptr_B - (idx_col_k * kElementsPerAccess - k) / kPackedElementsB,
true);
ElementAccumulator accum_SF_packed = ElementAccumulator(0);
for (int e = 0; e < kPackedElements; ++e) {
accum_SF_packed += A_converter(fragA[e]) * B_converter(fragB[e]);
}
accum_SF_packed *= SFA_converter(sfa) * SFB_converter(sfb);
accum += accum_SF_packed;
}
return accum;
}
__device__
void load_stages_gmem_to_smem(
int buffer_idx,
int num_stages,
int& unrolled_k,
int& global_k,
int tileA_k_local,
int smem_offset_A,
int smem_offset_B,
int smem_sf_write_offset,
bool is_even_thread,
bool load_b,
int K_limit,
ElementA const* ptr_A,
ElementB const* ptr_B,
ElementSFA const* ptr_SFA,
ElementSFB const* ptr_SFB,
SharedStorage& shared) {
for (int s = 0; s < num_stages; ++s) {
if (unrolled_k >= K_limit) {
return;
}
int sf_offset = compute_sf_offset_from_k(global_k);
bool stage_valid = (unrolled_k < K_limit);
if (is_even_thread) {
void* smem_ptr = &shared.smem_SFA[buffer_idx][s][smem_sf_write_offset];
ElementSFA const* gmem_ptr = ptr_SFA + sf_offset;
cutlass::arch::cp_async<sizeof(uint32_t)>(smem_ptr, gmem_ptr, stage_valid);
}
if (load_b && is_even_thread) {
int sf_sfb_offset = (threadIdx.x / 2) * 4;
void* smem_ptr = &shared.smem_SFB[buffer_idx][s][sf_sfb_offset];
ElementSFB const* gmem_ptr = ptr_SFB + sf_offset;
cutlass::arch::cp_async<sizeof(uint32_t)>(smem_ptr, gmem_ptr, stage_valid);
}
{
void* smem_ptr = &shared.smem_A[buffer_idx][s][smem_offset_A];
int packed_offset = unrolled_k / kPackedElementsA;
ElementA const* gmem_ptr = ptr_A + packed_offset;
cutlass::arch::cp_async<sizeof(FragmentA)>(smem_ptr, gmem_ptr, stage_valid);
}
if (load_b) {
void* smem_ptr = &shared.smem_B[buffer_idx][s][smem_offset_B];
int packed_offset = unrolled_k / kPackedElementsB;
ElementB const* gmem_ptr = ptr_B + packed_offset;
cutlass::arch::cp_async<sizeof(FragmentB)>(smem_ptr, gmem_ptr, stage_valid);
}
unrolled_k += tileA_k_local;
global_k += tileA_k_local;
}
}
__device__
void load_smem_fragments(
FragmentA& fragA,
FragmentB& fragB,
FragmentSFA& fragSFA,
FragmentSFB& fragSFB,
int smem_pipe_idx,
int k_block,
int smem_offset_A,
int smem_offset_B,
int smem_sf_offset,
SharedStorage& shared) {
cutlass::arch::shared_load(fragA, &shared.smem_A[smem_pipe_idx][k_block][smem_offset_A]);
cutlass::arch::shared_load(fragB, &shared.smem_B[smem_pipe_idx][k_block][smem_offset_B]);
{
uint32_t smem_ptr = cutlass::arch::cutlass_get_smem_pointer(
&shared.smem_SFA[smem_pipe_idx][k_block][smem_sf_offset]);
cutlass::arch::shared_load<2>(&fragSFA, smem_ptr);
}
{
int sfb_offset = threadIdx.x * kSFPerAccess;
uint32_t smem_ptr = cutlass::arch::cutlass_get_smem_pointer(
&shared.smem_SFB[smem_pipe_idx][k_block][sfb_offset]);
cutlass::arch::shared_load<2>(&fragSFB, smem_ptr);
}
}
__device__
ElementAccumulator warp_reduce_over_k(ElementAccumulator accum) {
for (int mask = (kThreadsPerRow >> 1); mask > 0; mask >>= 1) {
float other = __shfl_xor_sync(0xFFFFFFFFu, accum, mask, 32);
accum += other;
}
return accum;
}
__device__
void apply_epilogue(
ElementAccumulator accum,
ElementC* ptr_c_row)
{
*ptr_c_row = static_cast<ElementC>(__float2half(accum));
}
__global__
void nvfp4_gemv_kernel(Params params) {
extern __shared__ uint8_t smem_raw[];
SharedStorage& shared = *reinterpret_cast<SharedStorage*>(smem_raw);
int M = params.M;
int K = params.K;
int batch_count = params.batch_count;
for (int batch_idx = blockIdx.z; batch_idx < batch_count; batch_idx += gridDim.z) {
int idx_col_k = threadIdx.x;
int idx_row_m = blockIdx.x * blockDim.y + threadIdx.y;
if (idx_row_m >= M) {
continue;
}
ElementA const* ptr_A = params.ptr_A
+ batch_idx * params.batch_stride_A
+ idx_row_m * params.stride_A;
ptr_A += idx_col_k * (kElementsPerAccess / kPackedElementsA);
ElementB const* ptr_B = params.ptr_B
+ batch_idx * params.batch_stride_B;
ptr_B += idx_col_k * (kElementsPerAccess / kPackedElementsB);
ElementC* ptr_C_row = params.ptr_C
+ batch_idx * params.batch_stride_C
+ idx_row_m;
ElementSFA const* ptr_SFA = params.ptr_SFA + batch_idx * params.batch_stride_SFA;
ElementSFB const* ptr_SFB = params.ptr_SFB + batch_idx * params.batch_stride_SFB;
[[maybe_unused]] int SF_blocks_by_M = (M + 127) / 128;
int SF_blocks_by_K = (K / kSFVecSize + 3) / 4;
int row_block = idx_row_m >> 7;
int row_in_bloc = idx_row_m & 0x7f;
int row_low32 = row_in_bloc & 0x1f;
int row_hi = row_in_bloc >> 5;
ptr_SFA += (row_block * SF_blocks_by_K << 9)
+ (row_low32 << 4)
+ (row_hi << 2);
ElementAccumulator accum = ElementAccumulator(0);
const int tileA_k_local = kThreadsPerRow * kElementsPerAccess;
const int tiles_per_group = kStageCount;
const int macro_k = tileA_k_local * tiles_per_group;
const int full_groups = K / macro_k;
const int total_tiles = full_groups * tiles_per_group;
const int K_main = total_tiles * tileA_k_local;
int unrolled_k = 0;
int global_k = idx_col_k * kElementsPerAccess;
const int thread_id = threadIdx.y * kThreadsPerRow + threadIdx.x;
const bool is_even_thread = ((threadIdx.x % 2) == 0);
const bool load_b = (threadIdx.y == 0);
const int smem_sf_write_offset = (thread_id / 2) * 4;
const int smem_sf_offset = thread_id * kSFPerAccess;
if (total_tiles == 0) {
accum += process_tail_elements(
0,
idx_col_k,
K,
ptr_A,
ptr_B,
ptr_SFA,
ptr_SFB);
} else {
int smem_offset_A = thread_id * (kElementsPerAccess / kPackedElementsA);
int smem_offset_B = threadIdx.x * (kElementsPerAccess / kPackedElementsB);
for (int buf = 0; buf < kBufferCount - 1; ++buf) {
load_stages_gmem_to_smem(
buf,
kStageCount,
unrolled_k,
global_k,
tileA_k_local,
smem_offset_A,
smem_offset_B,
smem_sf_write_offset,
is_even_thread,
load_b,
K_main,
ptr_A,
ptr_B,
ptr_SFA,
ptr_SFB,
shared);
}
cutlass::arch::cp_async_fence();
cutlass::arch::cp_async_wait<kBufferCount - 2>();
__syncthreads();
FragmentA fragA_reg[2];
FragmentB fragB_reg[2];
FragmentSFA fragSFA_reg[2];
FragmentSFB fragSFB_reg[2];
int smem_pipe_read = 0;
int smem_pipe_write = kBufferCount - 1;
if (kStageCount > 1) {
int frag_idx = 0;
load_smem_fragments(
fragA_reg[frag_idx],
fragB_reg[frag_idx],
fragSFA_reg[frag_idx],
fragSFB_reg[frag_idx],
smem_pipe_read,
0,
smem_offset_A,
smem_offset_B,
smem_sf_offset,
shared);
}
int tile_idx = 0;
while (tile_idx < total_tiles) {
int smem_pipe_read_curr = smem_pipe_read;
for (int k_block = 0; k_block < kStageCount; ++k_block) {
if (k_block == kStageCount - 1) {
cutlass::arch::cp_async_wait<kBufferCount - 2>();
__syncthreads();
smem_pipe_read_curr = smem_pipe_read;
}
int k_block_next = (k_block + 1) % kStageCount;
int frag_idx_next = (k_block + 1) & 1;
load_smem_fragments(
fragA_reg[frag_idx_next],
fragB_reg[frag_idx_next],
fragSFA_reg[frag_idx_next],
fragSFB_reg[frag_idx_next],
smem_pipe_read_curr,
k_block_next,
smem_offset_A,
smem_offset_B,
smem_sf_offset,
shared);
if (k_block == 0) {
load_stages_gmem_to_smem(
smem_pipe_write,
kStageCount,
unrolled_k,
global_k,
tileA_k_local,
smem_offset_A,
smem_offset_B,
smem_sf_write_offset,
is_even_thread,
load_b,
K_main,
ptr_A,
ptr_B,
ptr_SFA,
ptr_SFB,
shared);
cutlass::arch::cp_async_fence();
smem_pipe_write = smem_pipe_read;
++smem_pipe_read;
if (smem_pipe_read == kBufferCount) {
smem_pipe_read = 0;
}
}
int frag_idx = k_block & 1;
accum += blockscaled_multiply_add(
fragA_reg[frag_idx],
fragB_reg[frag_idx],
fragSFA_reg[frag_idx],
fragSFB_reg[frag_idx]);
}
tile_idx += kStageCount;
}
cutlass::arch::cp_async_wait<0>();
__syncthreads();
if (unrolled_k < K) {
accum += process_tail_elements(
unrolled_k,
idx_col_k,
K,
ptr_A,
ptr_B,
ptr_SFA,
ptr_SFB);
}
}
accum = warp_reduce_over_k(accum);
if (threadIdx.x == 0) {
apply_epilogue(accum, ptr_C_row);
}
}
}
torch::Tensor nvfp4_gemv(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa_perm,
torch::Tensor sfb_perm,
torch::Tensor c)
{
TORCH_CHECK(a.is_cuda() && b.is_cuda() && c.is_cuda(),
"All tensors must be CUDA tensors");
int64_t M = a.size(0);
int64_t K_packed = a.size(1);
int64_t L = a.size(2);
int64_t K = K_packed * 2;
Params params;
params.M = static_cast<int>(M);
params.K = static_cast<int>(K);
params.batch_count = static_cast<int>(L);
auto a_strides = a.strides();
params.ptr_A = reinterpret_cast<ElementA const*>(a.data_ptr());
params.stride_A = a_strides[0];
params.batch_stride_A = a_strides[2];
auto b_strides = b.strides();
params.ptr_B = reinterpret_cast<ElementB const*>(b.data_ptr());
params.batch_stride_B = b_strides[2];
auto c_strides = c.strides();
params.ptr_C = reinterpret_cast<ElementC*>(c.data_ptr());
params.batch_stride_C = c_strides[2];
params.ptr_SFA = reinterpret_cast<ElementSFA const*>(sfa_perm.data_ptr());
params.ptr_SFB = reinterpret_cast<ElementSFB const*>(sfb_perm.data_ptr());
params.batch_stride_SFA = sfa_perm.numel() / L;
params.batch_stride_SFB = sfb_perm.numel() / L;
dim3 block(kThreadsPerRow, kThreadsPerCol, 1);
dim3 grid(
(M + block.y - 1) / block.y,
1,
params.batch_count);
size_t smem_size = sizeof(SharedStorage);
cudaStream_t stream = c10::cuda::getCurrentCUDAStream();
nvfp4_gemv_kernel<<<grid, block, smem_size, stream>>>(params);
return c;
}
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def(
"nvfp4_gemv",
&nvfp4_gemv_sketch::nvfp4_gemv,
"NVFP4 block-scaled GEMV (procedural sketch)");
}
"""
cpp_src = r"""
#include <torch/extension.h>
torch::Tensor nvfp4_gemv(
torch::Tensor a,
torch::Tensor b,
torch::Tensor sfa_perm,
torch::Tensor sfb_perm,
torch::Tensor c
);
"""
# Global cache for compiled extension
_nvfp4_ext = None
def get_extension():
"""Compile and cache the CUDA extension."""
global _nvfp4_ext
if _nvfp4_ext is not None:
return _nvfp4_ext
_nvfp4_ext = load_inline(
name="nvfp4_gemv_cutlass_probe",
cpp_sources=[cpp_src],
cuda_sources=[cuda_src],
extra_cflags=["-O3", "-std=c++17"],
extra_cuda_cflags=[
"-O3",
"-std=c++17",
"--expt-relaxed-constexpr",
"-use_fast_math",
"-gencode=arch=compute_100a,code=sm_100a"
],
extra_include_paths=CUTLASS_PATHS,
verbose=True,
)
return _nvfp4_ext
def custom_kernel(data: input_t) -> output_t:
"""
Skeleton implementation that zeros the output.
This validates the PyTorch <-> CUDA extension pipeline with CUTLASS headers.
"""
a, b, _, _, sfa_permuted, sfb_permuted, c = data
# Get compiled extension and call the kernel
ext = get_extension()
return ext.nvfp4_gemv(a, b, sfa_permuted, sfb_permuted, c)
scrolls · 647 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