submission 80995
Alex A · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 743 lines, June 9 Researcher Reciprocity License v1.0.
slow.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-80995?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:f22607401b4b1988e4a09077eea9835aa3a6ba968647b44639088c6774a536fc
license declaredunknown
license concludedunknown
authorsAlex A
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;fp4
using nvf4_t = cutlass::nv_float4_t<cutlass::float_e2m1_t>;fused-epilogue
using EpilogueTileShape = decltype(cute::make_shape(cute::Int<bM>{}, cute::Int<bN>{}));vector-width = uint4
uint4 packed[kTilesPerGroup];warp-specialization
cutlass::epilogue::TmaWarpSpecialized2Sm,Kernel source
slow.py743 lines
#!POPCORN leaderboard nvfp4_gemv
import os
import torch
from torch.utils.cpp_extension import load_inline
###############################################################################
# 1. C++ binding (pybind11) – exposes the Python entry point
###############################################################################
binding_cpp = r"""
#include <torch/extension.h>
void blockscaled_gemm_nvf4_f16(
at::Tensor& A,
at::Tensor& B,
at::Tensor& C,
at::Tensor& A_scales,
at::Tensor& B_scales,
int batch_size);
void to_blocked_and_gemm(
at::Tensor& A_fp4, at::Tensor& B_fp4,
at::Tensor& C,
at::Tensor& SFA, at::Tensor& SFB,
int batch_size);
void to_blocked_batched_cuda_wrapper(at::Tensor input, at::Tensor output);
void to_blocked_batched_dual_cuda_wrapper(
at::Tensor input0, at::Tensor output0,
at::Tensor input1, at::Tensor output1);
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("blockscaled_gemm_nvf4_f16",
&blockscaled_gemm_nvf4_f16,
"NVF4 blockscaled GEMV with FP16 output");
m.def("to_blocked_batched_cuda",
&to_blocked_batched_cuda_wrapper,
"Batched to_blocked conversion (CUDA)");
m.def("to_blocked_batched_dual_cuda",
&to_blocked_batched_dual_cuda_wrapper,
"Dual batched to_blocked conversion (CUDA)");
m.def("to_blocked_and_gemm",
&to_blocked_and_gemm,
"Run scale blocking followed by blockscaled GEMM");
}
"""
###############################################################################
# 2. CUDA / CUTLASS kernel implementation (inline translation unit)
###############################################################################
cuda_src = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <ATen/ATen.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDACachingAllocator.h>
#include <c10/cuda/CUDAGuard.h>
#include <torch/extension.h>
#include <iostream>
#include <type_traits>
#include <cstdlib>
#include <cstdint>
#include <limits>
#include <cutlass/cutlass.h>
#include <cutlass/numeric_types.h>
#include <cutlass/epilogue/fusion/operations.hpp>
#include <cutlass/epilogue/thread/linear_combination.h>
#include <cutlass/epilogue/collective/collective_builder.hpp>
#include <cutlass/gemm/collective/collective_builder.hpp>
#include <cutlass/gemm/device/gemm_universal_adapter.h>
#include <cutlass/gemm/kernel/gemm_universal.hpp>
#include <cutlass/gemm/kernel/tile_scheduler_params.h>
#include <cutlass/detail/sm100_blockscaled_layout.hpp>
#include <cutlass/util/packed_stride.hpp>
#include <cute/tensor.hpp>
// ============================================================================
// CUDA kernel for fast batched to_blocked (supports two tensors per launch)
// ============================================================================
struct TensorParams {
const uint8_t* input;
uint8_t* output;
int rows;
int cols;
int batches;
int n_row_blocks;
int n_col_blocks;
int n_col_groups;
int64_t in_row_stride;
int64_t in_col_stride;
int64_t in_batch_stride;
int64_t out_batch_stride;
};
constexpr int kWarpSize = 32;
constexpr int kTilesPerCTA = 4;
constexpr int kTilesPerGroup = 4;
__global__ void to_blocked_batched_kernel(
const TensorParams params0,
const TensorParams params1,
size_t total_tiles0,
size_t total_tiles1) {
const int warp_idx = threadIdx.y;
const int lane = threadIdx.x;
const size_t tile_linear =
static_cast<size_t>(blockIdx.x) * blockDim.y + warp_idx;
const size_t total_tiles = total_tiles0 + total_tiles1;
if (tile_linear >= total_tiles) {
return;
}
const bool use_first = (tile_linear < total_tiles0);
const TensorParams params =
use_first ? params0 : params1;
const size_t local_tile =
use_first ? tile_linear : (tile_linear - total_tiles0);
const int rows = params.rows;
const int batches = params.batches;
const int n_row_blocks = params.n_row_blocks;
const int n_col_blocks = params.n_col_blocks;
const int n_col_groups = params.n_col_groups;
const int64_t in_batch_stride = params.in_batch_stride;
const int64_t in_row_stride = params.in_row_stride;
const int64_t out_batch_stride= params.out_batch_stride;
const uint8_t* input = params.input;
uint8_t* output = params.output;
const size_t tile_groups_per_batch =
static_cast<size_t>(n_row_blocks) * n_col_groups;
if (tile_groups_per_batch == 0) {
return;
}
const int batch = static_cast<int>(local_tile / tile_groups_per_batch);
if (batch >= batches) {
return;
}
const size_t tile_in_batch = local_tile % tile_groups_per_batch;
const int rb = static_cast<int>(tile_in_batch / n_col_groups);
const int cb_group = static_cast<int>(tile_in_batch % n_col_groups);
const int cb_base = cb_group * kTilesPerGroup;
int tiles_in_group = n_col_blocks - cb_base;
if (tiles_in_group <= 0) {
return;
}
if (tiles_in_group > kTilesPerGroup) {
tiles_in_group = kTilesPerGroup;
}
const int tile_row_base = rb * 128;
if (tile_row_base >= rows) {
return;
}
int row_limit = rows - tile_row_base;
if (row_limit > 128) {
row_limit = 128;
}
const size_t tile_idx_base =
static_cast<size_t>(rb) * n_col_blocks + cb_base;
const size_t tile_base_base =
static_cast<size_t>(batch) * out_batch_stride + tile_idx_base * 512;
// Precompute the base pointer to the first row of this row-block & batch.
const int64_t col_offset_bytes = static_cast<int64_t>(cb_base) * 4;
const uint8_t* base_row_ptr =
input +
static_cast<int64_t>(batch) * in_batch_stride +
static_cast<int64_t>(tile_row_base) * in_row_stride +
col_offset_bytes;
uint4 packed[kTilesPerGroup];
#pragma unroll
for (int t = 0; t < kTilesPerGroup; ++t) {
packed[t] = make_uint4(0u, 0u, 0u, 0u);
}
// Each warp covers 128 rows (4 groups * 32 lanes).
#pragma unroll
for (int group = 0; group < 4; ++group) {
const int logical_row = lane + group * kWarpSize;
const bool valid = (logical_row < row_limit);
const uint8_t* row_ptr =
base_row_ptr + static_cast<int64_t>(logical_row) * in_row_stride;
// Masked 128b load via read-only path.
uint4 chunk = valid
? __ldg(reinterpret_cast<const uint4*>(row_ptr))
: make_uint4(0u, 0u, 0u, 0u);
const uint32_t w0 = chunk.x;
const uint32_t w1 = chunk.y;
const uint32_t w2 = chunk.z;
const uint32_t w3 = chunk.w;
if (tiles_in_group > 0) reinterpret_cast<uint32_t*>(&packed[0])[group] = w0;
if (tiles_in_group > 1) reinterpret_cast<uint32_t*>(&packed[1])[group] = w1;
if (tiles_in_group > 2) reinterpret_cast<uint32_t*>(&packed[2])[group] = w2;
if (tiles_in_group > 3) reinterpret_cast<uint32_t*>(&packed[3])[group] = w3;
}
#pragma unroll
for (int t = 0; t < kTilesPerGroup; ++t) {
if (t >= tiles_in_group) {
break;
}
const size_t tile_base =
tile_base_base + static_cast<size_t>(t) * 512;
uint8_t* tile_out_base =
output + tile_base + static_cast<size_t>(lane) * 16;
// Streaming 128-bit store: no L1 allocation, minimal L2 pollution.
uint4 v = packed[t];
asm volatile(
"st.global.cs.v4.b32 [%0], {%1,%2,%3,%4};\n"
:
: "l"(tile_out_base),
"r"(v.x), "r"(v.y), "r"(v.z), "r"(v.w)
);
}
}
namespace {
TensorParams make_params(const at::Tensor& input, const at::Tensor& output) {
TensorParams params;
params.input = input.data_ptr<uint8_t>();
params.output = output.data_ptr<uint8_t>();
params.rows = static_cast<int>(input.size(0));
params.cols = static_cast<int>(input.size(1));
params.batches = static_cast<int>(input.size(2));
params.n_row_blocks = (params.rows + 127) / 128;
params.n_col_blocks = (params.cols + 3) / 4;
params.n_col_groups = (params.n_col_blocks + kTilesPerGroup - 1) / kTilesPerGroup;
params.in_row_stride = input.stride(0);
params.in_col_stride = input.stride(1);
params.in_batch_stride = input.stride(2);
params.out_batch_stride = static_cast<int64_t>(params.rows) * params.cols;
return params;
}
size_t count_tile_groups(const TensorParams& params) {
return static_cast<size_t>(params.batches) *
params.n_row_blocks *
params.n_col_groups;
}
void launch_to_blocked(const TensorParams& params0, size_t groups0,
const TensorParams& params1, size_t groups1) {
size_t total_groups = groups0 + groups1;
if (total_groups == 0) {
return;
}
TORCH_CHECK(total_groups <= static_cast<size_t>(std::numeric_limits<int>::max()),
"Total tile group count exceeds supported range");
auto validate = [](const TensorParams& params, size_t groups) {
if (groups == 0) {
return;
}
TORCH_CHECK(params.in_col_stride == 1,
"Expected unit column stride for blocked conversion");
TORCH_CHECK((reinterpret_cast<uintptr_t>(params.input) & 0xF) == 0,
"Input tensor must be 16-byte aligned");
TORCH_CHECK((reinterpret_cast<uintptr_t>(params.output) & 0xF) == 0,
"Output tensor must be 16-byte aligned");
};
validate(params0, groups0);
validate(params1, groups1);
auto stream = at::cuda::getCurrentCUDAStream();
unsigned int grid_blocks = static_cast<unsigned int>((total_groups + kTilesPerCTA - 1) / kTilesPerCTA);
dim3 grid(grid_blocks);
dim3 block(kWarpSize, kTilesPerCTA);
to_blocked_batched_kernel<<<grid, block, 0, stream>>>(
params0, params1, groups0, groups1);
}
} // namespace
void to_blocked_batched_cuda_wrapper(at::Tensor input, at::Tensor output) {
int rows = input.size(0);
int cols = input.size(1);
TORCH_CHECK(rows % 128 == 0, "to_blocked_batched expects rows to be a multiple of 128");
TORCH_CHECK(cols % 4 == 0, "to_blocked_batched expects cols to be a multiple of 4");
TensorParams params0 = make_params(input, output);
size_t tiles0 = count_tile_groups(params0);
TensorParams empty{};
launch_to_blocked(params0, tiles0, empty, 0);
}
void to_blocked_batched_dual_cuda_wrapper(
at::Tensor input0, at::Tensor output0,
at::Tensor input1, at::Tensor output1) {
int rows0 = input0.size(0);
int cols0 = input0.size(1);
int rows1 = input1.size(0);
int cols1 = input1.size(1);
TORCH_CHECK(rows0 % 128 == 0, "A_scales rows must be multiple of 128");
TORCH_CHECK(cols0 % 4 == 0, "A_scales cols must be multiple of 4");
TORCH_CHECK(rows1 % 128 == 0, "B_scales rows must be multiple of 128");
TORCH_CHECK(cols1 % 4 == 0, "B_scales cols must be multiple of 4");
TensorParams params0 = make_params(input0, output0);
TensorParams params1 = make_params(input1, output1);
size_t tiles0 = count_tile_groups(params0);
size_t tiles1 = count_tile_groups(params1);
launch_to_blocked(params0, tiles0, params1, tiles1);
}
#define CUTLASS_CHECK(status) \
{ \
cutlass::Status error = status; \
if (error != cutlass::Status::kSuccess) { \
std::cerr << "Got CUTLASS error: " << cutlassGetStatusString(error) << " at line " \
<< __LINE__ << std::endl; \
exit(EXIT_FAILURE); \
} \
}
struct GemmParams {
int m, n, k, l;
void *ptr_A, *ptr_B, *ptr_C;
void *ptr_SFA, *ptr_SFB;
float alpha, beta;
};
using nvf4_t = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
template <typename ElementA, typename ElementB, typename ElementC,
typename ElementAccum, int bM_, int bN_, int bK_>
struct CollectiveTemplates {
using LayoutATag = cutlass::layout::RowMajor;
using LayoutBTag = cutlass::layout::ColumnMajor;
using LayoutCTag = cutlass::layout::RowMajor;
static constexpr int AlignmentA = cute::is_same_v<ElementA, nvf4_t> ? 128 : 16;
static constexpr int AlignmentB = cute::is_same_v<ElementB, nvf4_t> ? 128 : 16;
static constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
using ElementAccumulator = ElementAccum;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
static constexpr int bM = bM_;
static constexpr int bN = bN_;
static constexpr int bK = bK_;
using MmaTileShape = decltype(cute::make_shape(cute::Int<bM>{}, cute::Int<bN>{}, cute::Int<bK>{}));
using EpilogueTileShape = decltype(cute::make_shape(cute::Int<bM>{}, cute::Int<bN>{}));
using ClusterShape = cute::Shape<cute::_2, cute::_1, cute::_1>;
using ProblemShape = cute::Shape<int, int, int, int>;
using DefaultEpilogueOp = cutlass::epilogue::fusion::LinearCombination<ElementC, ElementAccum>;
using DefaultCollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag,
OperatorClass,
MmaTileShape,
ClusterShape,
cutlass::epilogue::collective::EpilogueTileAuto,
ElementAccumulator,
ElementAccumulator,
void,
LayoutCTag,
AlignmentC,
ElementC,
LayoutCTag,
AlignmentC,
cutlass::epilogue::TmaWarpSpecialized2Sm,
DefaultEpilogueOp
>::CollectiveOp;
template <typename CollectiveEpilogue_>
using DefaultCollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
ArchTag,
OperatorClass,
ElementA,
LayoutATag,
AlignmentA,
ElementB,
LayoutBTag,
AlignmentB,
ElementAccumulator,
MmaTileShape,
ClusterShape,
cutlass::gemm::collective::StageCountAutoCarveout<
static_cast<int>(sizeof(typename CollectiveEpilogue_::SharedStorage))>,
cutlass::gemm::KernelTmaWarpSpecialized2SmBlockScaledSm100
>::CollectiveOp;
template <typename CollectiveMainloop_, typename CollectiveEpilogue_>
using DefaultGemmKernel = cutlass::gemm::kernel::GemmUniversal<
ProblemShape,
CollectiveMainloop_,
CollectiveEpilogue_,
void>;
};
template <int bM_, int bN_, int bK_,
typename ElementA, typename ElementB,
typename ElementC, typename ElementAccum>
void run_blockscaled_gemm(GemmParams& params, cudaStream_t stream) {
using Collectives = CollectiveTemplates<ElementA, ElementB, ElementC, ElementAccum, bM_, bN_, bK_>;
using CollectiveEpilogue = typename Collectives::DefaultCollectiveEpilogue;
using CollectiveMainloop = typename Collectives::template DefaultCollectiveMainloop<CollectiveEpilogue>;
using GemmKernel = typename Collectives::template DefaultGemmKernel<CollectiveMainloop, CollectiveEpilogue>;
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
using StrideA = typename Gemm::GemmKernel::StrideA;
using StrideB = typename Gemm::GemmKernel::StrideB;
using StrideC = typename Gemm::GemmKernel::StrideC;
using StrideD = typename Gemm::GemmKernel::StrideD;
StrideA stride_A = cutlass::make_cute_packed_stride(StrideA{}, {params.m, params.k, params.l});
StrideB stride_B = cutlass::make_cute_packed_stride(StrideB{}, {params.n, params.k, params.l});
StrideC stride_C = cutlass::make_cute_packed_stride(StrideC{}, {params.m, params.n, params.l});
StrideD stride_D = stride_C;
auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(
cute::make_shape(params.m, params.n, params.k, params.l));
auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(
cute::make_shape(params.m, params.n, params.k, params.l));
cutlass::KernelHardwareInfo hw_info;
hw_info.device_id = 0;
hw_info.sm_count = cutlass::KernelHardwareInfo::query_device_multiprocessor_count(hw_info.device_id);
hw_info.cluster_shape = dim3(2, 1, 1);
hw_info.cluster_shape_fallback = dim3(2, 1, 1);
typename Gemm::GemmKernel::TileSchedulerArguments scheduler;
typename Gemm::Arguments arguments;
decltype(arguments.epilogue.thread) fusion_args;
fusion_args.alpha = params.alpha;
fusion_args.beta = params.beta;
arguments = typename Gemm::Arguments{
cutlass::gemm::GemmUniversalMode::kGemm,
{params.m, params.n, params.k, params.l},
{
static_cast<typename ElementA::DataType*>(params.ptr_A), stride_A,
static_cast<typename ElementB::DataType*>(params.ptr_B), stride_B,
static_cast<typename ElementA::ScaleFactorType*>(params.ptr_SFA), layout_SFA,
static_cast<typename ElementB::ScaleFactorType*>(params.ptr_SFB), layout_SFB
},
{
fusion_args,
nullptr, stride_C,
static_cast<ElementC*>(params.ptr_C), stride_D
},
hw_info,
scheduler
};
//
// -------- FIX #1: Persistent workspace -------------
//
size_t workspace_size = Gemm::get_workspace_size(arguments);
static thread_local void* workspace_ptr = nullptr;
static thread_local size_t workspace_cap = 0;
if (workspace_size > workspace_cap) {
if (workspace_ptr) {
c10::cuda::CUDACachingAllocator::raw_delete(workspace_ptr);
}
workspace_ptr = c10::cuda::CUDACachingAllocator::raw_alloc_with_stream(
workspace_size, stream);
workspace_cap = workspace_size;
}
//
// -------- FIX #2: Persistent initialized GEMM object -------------
//
static thread_local Gemm gemm;
static thread_local bool initialized = false;
if (!initialized) {
CUTLASS_CHECK(gemm.can_implement(arguments));
CUTLASS_CHECK(gemm.initialize(arguments, workspace_ptr, stream));
initialized = true;
} else {
// Only update problem sizes + pointers
CUTLASS_CHECK(gemm.update(arguments));
}
//
// -------- Execute GEMM -------------
//
CUTLASS_CHECK(gemm.run(stream));
}
template void run_blockscaled_gemm<256, 64, 256, nvf4_t, nvf4_t, cutlass::half_t, float>(
GemmParams&, cudaStream_t);
void blockscaled_gemm_nvf4_f16(
at::Tensor& A,
at::Tensor& B,
at::Tensor& C,
at::Tensor& A_scales,
at::Tensor& B_scales,
int batch_size) {
TORCH_CHECK(A.is_cuda(), "A must be CUDA");
TORCH_CHECK(B.is_cuda(), "B must be CUDA");
TORCH_CHECK(C.is_cuda(), "C must be CUDA");
TORCH_CHECK(A_scales.is_cuda(), "A_scales must be CUDA");
TORCH_CHECK(B_scales.is_cuda(), "B_scales must be CUDA");
GemmParams params;
params.m = static_cast<int>(A.size(0));
params.n = static_cast<int>(B.size(0));
params.k = static_cast<int>(A.size(1)) * 2;
params.l = static_cast<int>(A.size(2));
params.ptr_A = A.data_ptr();
params.ptr_B = B.data_ptr();
params.ptr_C = C.data_ptr();
params.ptr_SFA = A_scales.data_ptr();
params.ptr_SFB = B_scales.data_ptr();
params.alpha = 1.0f;
params.beta = 0.0f;
auto stream = at::cuda::getCurrentCUDAStream().stream();
run_blockscaled_gemm<256, 64, 256, nvf4_t, nvf4_t, cutlass::half_t, float>(params, stream);
}
at::Tensor _make_blocked_tensor_like(const at::Tensor& src) {
auto options = src.options().dtype(at::kByte);
auto rows = src.size(0);
auto cols = src.size(1);
auto batches = src.size(2);
TORCH_CHECK(src.dim() == 3, "Expected scale tensor with 3 dimensions");
auto total_elements = rows * cols * batches;
auto buffer = at::empty({total_elements}, options);
auto blocked = buffer.as_strided(
{rows, cols, batches},
{cols, 1, rows * cols});
return blocked;
}
void to_blocked_and_gemm(
at::Tensor& A_fp4, at::Tensor& B_fp4,
at::Tensor& C,
at::Tensor& SFA, at::Tensor& SFB,
int batch_size) {
TORCH_CHECK(A_fp4.is_cuda(), "A_fp4 must be CUDA");
TORCH_CHECK(B_fp4.is_cuda(), "B_fp4 must be CUDA");
TORCH_CHECK(C.is_cuda(), "C must be CUDA");
TORCH_CHECK(SFA.is_cuda(), "SFA must be CUDA");
TORCH_CHECK(SFB.is_cuda(), "SFB must be CUDA");
TORCH_CHECK(A_fp4.dim() == 3, "A tensor must be [m, k_packed, l]");
TORCH_CHECK(B_fp4.dim() == 3, "B tensor must be [n, k_packed, l]");
TORCH_CHECK(C.dim() == 3, "C tensor must be [m, n, l]");
TORCH_CHECK(SFA.size(2) == batch_size, "SFA batch dimension mismatch");
TORCH_CHECK(SFB.size(2) == batch_size, "SFB batch dimension mismatch");
c10::cuda::OptionalCUDAGuard guard(A_fp4.device());
auto torch_stream = at::cuda::getCurrentCUDAStream();
auto cuda_stream = torch_stream.stream();
auto A_scales_blocked = _make_blocked_tensor_like(SFA);
auto B_scales_blocked = _make_blocked_tensor_like(SFB);
to_blocked_batched_dual_cuda_wrapper(SFA, A_scales_blocked, SFB, B_scales_blocked);
GemmParams params;
params.m = static_cast<int>(A_fp4.size(0));
params.n = static_cast<int>(B_fp4.size(0));
params.k = static_cast<int>(A_fp4.size(1)) * 2;
params.l = batch_size;
params.ptr_A = A_fp4.data_ptr();
params.ptr_B = B_fp4.data_ptr();
params.ptr_C = C.data_ptr();
params.ptr_SFA = A_scales_blocked.data_ptr();
params.ptr_SFB = B_scales_blocked.data_ptr();
params.alpha = 1.0f;
params.beta = 0.0f;
run_blockscaled_gemm<256, 64, 256, nvf4_t, nvf4_t, cutlass::half_t, float>(params, cuda_stream);
}
"""
###############################################################################
# 3. Build the extension
###############################################################################
_THIS_DIR = os.path.dirname(__file__)
_INCLUDE_DIRS = [
"/usr/local/include",
"/usr/include",
"/opt/cutlass/include",
os.path.join(_THIS_DIR, "cutlass", "include"),
os.path.abspath(os.path.join(_THIS_DIR, "..", "..", "..", "..", "cutlass", "include")),
os.path.join(_THIS_DIR, "cutlass", "tools", "util", "include"),
os.path.abspath(os.path.join(_THIS_DIR, "..", "..", "..", "..", "cutlass", "tools", "util", "include")),
]
_CUDA_FLAGS = [
"-O3",
"--use_fast_math",
"-std=c++17",
"-gencode=arch=compute_100a,code=sm_100a",
]
_CUDA_FLAGS.extend(f"-I{path}" for path in _INCLUDE_DIRS)
ext = load_inline(
name="fp4gemv_ext_v9",
cpp_sources=[binding_cpp],
cuda_sources=[cuda_src],
extra_cflags=["-O3", "-std=c++17"],
extra_cuda_cflags=_CUDA_FLAGS,
verbose=False,
)
###############################################################################
# 4. Helper utilities shared with the kernel wrapper
###############################################################################
def ensure_packed_fp4(tensor: torch.Tensor, rows: int, k_elements: int, batches: int) -> torch.Tensor:
"""
Ensure FP4 tensor is packed as NVF4 bytes (two 4-bit values per byte).
Returns a tensor of shape [rows, k_elements//2, batches] in uint8.
"""
bytes_view = tensor.view(torch.uint8).reshape(rows, -1, batches)
bytes_per_row = bytes_view.size(1)
target_bytes = k_elements // 2
if bytes_per_row != target_bytes:
raise RuntimeError(f"Expected packed NVF4 with {target_bytes} bytes/row, got {bytes_per_row}")
return bytes_view
def _allocate_blocked_buffer(rows: int, cols: int, batches: int, device):
buf = torch.empty((rows * cols * batches,), dtype=torch.uint8, device=device)
return buf.as_strided((rows, cols, batches), (cols, 1, rows * cols))
def to_blocked_batched_pair(t0: torch.Tensor, t1: torch.Tensor):
r0, c0, b0 = t0.shape
r1, c1, b1 = t1.shape
out0 = _allocate_blocked_buffer(r0, c0, b0, t0.device)
out1 = _allocate_blocked_buffer(r1, c1, b1, t1.device)
ext.to_blocked_batched_dual_cuda(t0, out0, t1, out1)
return out0, out1
def to_blocked_batched_fast(t: torch.Tensor) -> torch.Tensor:
r, c, b = t.shape
out = _allocate_blocked_buffer(r, c, b, t.device)
ext.to_blocked_batched_cuda(t, out)
return out
def _require_profiler_stride(tensor: torch.Tensor, dim0: int, dim1: int, batches: int, name: str):
expected_stride = (dim1, 1, dim0 * dim1)
if tensor.stride() != expected_stride:
raise RuntimeError(
f"{name} must use profiler stride {expected_stride}, got {tensor.stride()}"
)
def blockscaled_gemm(A: torch.Tensor,
B: torch.Tensor,
C: torch.Tensor,
A_scales: torch.Tensor,
B_scales: torch.Tensor,
batch_size: int = 1) -> torch.Tensor:
m, k_packed, l = A.shape
n = B.size(0)
_require_profiler_stride(A, m, k_packed, l, "A")
_require_profiler_stride(B, n, k_packed, l, "B")
_require_profiler_stride(C, m, n, l, "C")
ext.to_blocked_and_gemm(A, B, C, A_scales, B_scales, batch_size)
return C
###############################################################################
# 5. GPU submission entry point expected by the leaderboard harness
###############################################################################
def custom_kernel(submission_input):
a_ref, b_ref, sfa_ref, sfb_ref, _, _, c_template = submission_input
m, _, l = a_ref.shape
tile_n = b_ref.size(0)
# Create C with profiler-style stride to avoid conversion
C_buf = torch.empty((m * tile_n * l,), dtype=torch.float16, device=c_template.device)
C_full = C_buf.as_strided((m, tile_n, l), (tile_n, 1, m * tile_n))
sfa_u8 = sfa_ref.view(torch.uint8)
sfb_u8 = sfb_ref.view(torch.uint8)
blockscaled_gemm(
a_ref,
b_ref,
C_full,
sfa_u8,
sfb_u8,
batch_size=l,
)
return C_full.narrow(1, 0, 1)
scrolls · 743 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