submission 97474
CatsRCool · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 368 lines, June 9 Researcher Reciprocity License v1.0.
s_attempt.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemv-97474?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:f367021859a005605a780c61945a9ba38e144aeb6750bdeae713fb23e1a754b1
license declaredunknown
license concludedunknown
authorsCatsRCool
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
template <typename ClusterShape>fp4
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;fused-epilogue
using FusionOperation = cutlass::epilogue::fusion::LinearCombination<ElementD, ElementCompute>;Kernel source
s_attempt.py368 lines
import torch
from task import input_t, output_t
from torch.utils.cpp_extension import load_inline
import cutlass
import cutlass.cute as cute
from cutlass.cute.runtime import make_ptr
import cutlass.utils.blockscaled_layout as blockscaled_utils
# ---------------------------------------------------------------------------
# CUTLASS Tensor Core GEMV - Optimized with inline column extraction
# ---------------------------------------------------------------------------
cuda_source = """
#include <cuda_fp16.h>
#include <cuda_fp8.h>
#include <cuda_runtime.h>
#include <stdexcept>
#include <map>
#include <tuple>
#include <cstdint>
// Prevent CUTLASS synclog from being included, then stub out the functions
#define CUTLASS_ARCH_SYNCLOG_H_
namespace cutlass {
namespace arch {
template<typename... Args>
__device__ __host__ inline void synclog_emit_tma_load(Args...) {}
template<typename... Args>
__device__ __host__ inline void synclog_emit_tma_store(Args...) {}
template<typename... Args>
__device__ __host__ inline void synclog_emit_fence_view_async_shared(Args...) {}
template<typename... Args>
__device__ __host__ inline void synclog_emit_tma_store_arrive(Args...) {}
template<typename... Args>
__device__ __host__ inline void synclog_emit_tma_store_wait(Args...) {}
}
}
#include "cutlass/cutlass.h"
#include <ATen/cuda/CUDAContext.h>
#include "cute/tensor.hpp"
#include "cutlass/gemm/collective/collective_builder.hpp"
#include "cutlass/epilogue/collective/collective_builder.hpp"
#include "cutlass/detail/sm100_blockscaled_layout.hpp"
#include "cutlass/gemm/device/gemm_universal_adapter.h"
#include "cutlass/gemm/kernel/gemm_universal.hpp"
#include "cutlass/util/packed_stride.hpp"
#include "cutlass/util/device_memory.h"
#include "cutlass/cuda_host_adapter.hpp"
using namespace cute;
// ---------------------------------------------------------------------------
// Optimized Column Extraction Kernel (vectorized loads/stores)
// ---------------------------------------------------------------------------
__global__ void extract_col0_vectorized_kernel(
const __half* __restrict__ gemm_output, // [m, n=32, l] RowMajor layout
__half* __restrict__ gemv_output, // [m, l]
int m, int n, int l
) {
// Thread handles one element at a time
int idx = blockIdx.x * blockDim.x + threadIdx.x;
int total = m * l;
if (idx < total) {
// Decode output position: idx = batch * m + row
int batch = idx / m;
int row = idx % m;
// CUTLASS LayoutD is RowMajor, so stride: (n, 1, m*n)
// Index for (row, col=0, batch): row*n + 0 + batch*m*n
int gemm_idx = row * n + batch * m * n;
gemv_output[idx] = gemm_output[gemm_idx];
}
}
// ---------------------------------------------------------------------------
// Global Configurations
// ---------------------------------------------------------------------------
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutATag = cutlass::layout::RowMajor;
constexpr int AlignmentA = 32;
using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutBTag = cutlass::layout::ColumnMajor;
constexpr int AlignmentB = 32;
using ElementC = cutlass::half_t;
using ElementD = cutlass::half_t;
using LayoutCTag = cutlass::layout::RowMajor;
using LayoutDTag = cutlass::layout::RowMajor; // RowMajor for easier column extraction
constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
using ElementAccumulator = float;
using ElementCompute = float;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
using MmaTileShape = Shape<_128, _64, _256>; // M=128 is minimum, try larger K for better arithmetic intensity
using FusionOperation = cutlass::epilogue::fusion::LinearCombination<ElementD, ElementCompute>;
// ---------------------------------------------------------------------------
// Templated Gemm Config
// ---------------------------------------------------------------------------
template <typename ClusterShape>
struct GemmConfig {
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag, OperatorClass,
MmaTileShape, ClusterShape,
cutlass::epilogue::collective::EpilogueTileAuto,
ElementAccumulator, ElementAccumulator,
ElementC, LayoutCTag, AlignmentC,
ElementD, LayoutDTag, AlignmentD,
cutlass::epilogue::collective::EpilogueScheduleAuto,
FusionOperation
>::CollectiveOp;
using CollectiveMainloop = 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::collective::KernelScheduleAuto
>::CollectiveOp;
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
Shape<int, int, int, int>,
CollectiveMainloop,
CollectiveEpilogue,
void>;
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
};
// ---------------------------------------------------------------------------
// Helper to run GEMM
// ---------------------------------------------------------------------------
static std::map<int, torch::Tensor> workspace_cache;
template <typename Config>
void run_gemm(
torch::Tensor A, torch::Tensor B, torch::Tensor SFA, torch::Tensor SFB,
torch::Tensor C_out,
int m, int n_problem, int k, int l, int n_phys_padded
) {
using Gemm = typename Config::Gemm;
using StrideA = typename Gemm::GemmKernel::StrideA;
using StrideB = typename Gemm::GemmKernel::StrideB;
using StrideC = typename Gemm::GemmKernel::StrideC;
using StrideD = typename Gemm::GemmKernel::StrideD;
using LayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFA;
using LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFB;
using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
StrideA stride_A = cutlass::make_cute_packed_stride(StrideA{}, {m, k, l});
StrideB stride_B = cutlass::make_cute_packed_stride(StrideB{}, {n_phys_padded, k, l});
StrideC stride_C = cutlass::make_cute_packed_stride(StrideC{}, {m, n_problem, l});
StrideD stride_D = cutlass::make_cute_packed_stride(StrideD{}, {m, n_problem, l});
LayoutSFA layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(m, n_phys_padded, k, l));
LayoutSFB layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(m, n_phys_padded, k, l));
typename Gemm::Arguments arguments{
cutlass::gemm::GemmUniversalMode::kGemm,
{m, n_problem, k, l},
{
reinterpret_cast<ElementA::DataType*>(A.data_ptr()), stride_A,
reinterpret_cast<ElementB::DataType*>(B.data_ptr()), stride_B,
reinterpret_cast<ElementA::ScaleFactorType*>(SFA.data_ptr()), layout_SFA,
reinterpret_cast<ElementB::ScaleFactorType*>(SFB.data_ptr()), layout_SFB
},
{
{ElementCompute(1.0f), ElementCompute(0.0f)},
nullptr, stride_C,
reinterpret_cast<ElementD*>(C_out.data_ptr()), stride_D
}
};
Gemm gemm;
size_t workspace_size = Gemm::get_workspace_size(arguments);
int device_id = A.device().index();
void* workspace_ptr = nullptr;
if (workspace_size > 0) {
auto ws_it = workspace_cache.find(device_id);
if (ws_it == workspace_cache.end() || ws_it->second.numel() * ws_it->second.element_size() < workspace_size) {
torch::Tensor workspace_tensor = torch::empty({(int64_t)workspace_size},
torch::TensorOptions().dtype(torch::kUInt8).device(A.device()));
workspace_cache[device_id] = workspace_tensor;
workspace_ptr = workspace_tensor.data_ptr();
} else {
workspace_ptr = ws_it->second.data_ptr();
}
}
cutlass::Status status = gemm.can_implement(arguments);
if (status != cutlass::Status::kSuccess) throw std::runtime_error("CUTLASS kernel cannot implement");
status = gemm.initialize(arguments, workspace_ptr);
if (status != cutlass::Status::kSuccess) throw std::runtime_error("CUTLASS kernel init failed");
status = gemm.run(at::cuda::getCurrentCUDAStream());
if (status != cutlass::Status::kSuccess) throw std::runtime_error("CUTLASS kernel run failed");
}
// ---------------------------------------------------------------------------
// Temporary output cache
// ---------------------------------------------------------------------------
static std::map<std::tuple<int, int, int>, torch::Tensor> temp_output_cache;
void fp4_gemv_cutlass(
torch::Tensor A,
torch::Tensor B,
torch::Tensor SFA,
torch::Tensor SFB,
torch::Tensor C_out,
int m, int k, int l,
int sfa_s0, int sfa_s1, int sfa_s2, int sfa_s3, int sfa_s4, int sfa_s5,
int sfb_s0, int sfb_s1, int sfb_s2, int sfb_s3, int sfb_s4, int sfb_s5
) {
constexpr int n_phys_padded = 128;
constexpr int n_problem = 32;
// Get or allocate cached temporary output buffer [m, n_problem, l]
int device_id = A.device().index();
auto temp_key = std::make_tuple(m, l, device_id);
torch::Tensor temp_output_tensor;
auto temp_it = temp_output_cache.find(temp_key);
if (temp_it == temp_output_cache.end()) {
temp_output_tensor = torch::empty({m, n_problem, l},
torch::TensorOptions().dtype(torch::kFloat16).device(A.device()));
temp_output_cache[temp_key] = temp_output_tensor;
} else {
temp_output_tensor = temp_it->second;
}
// Run GEMM to temporary buffer
if (m >= 512) {
run_gemm<GemmConfig<Shape<_4, _1, _1>>>(A, B, SFA, SFB, temp_output_tensor, m, n_problem, k, l, n_phys_padded);
} else {
run_gemm<GemmConfig<Shape<_1, _1, _1>>>(A, B, SFA, SFB, temp_output_tensor, m, n_problem, k, l, n_phys_padded);
}
// Extract column 0 using optimized kernel
cudaStream_t stream = at::cuda::getCurrentCUDAStream();
int total_elements = m * l;
int threads_per_block = 256;
int num_blocks = (total_elements + threads_per_block - 1) / threads_per_block;
extract_col0_vectorized_kernel<<<num_blocks, threads_per_block, 0, stream>>>(
reinterpret_cast<__half*>(temp_output_tensor.data_ptr()),
reinterpret_cast<__half*>(C_out.data_ptr()),
m, n_problem, l
);
cudaError_t err = cudaGetLastError();
if (err != cudaSuccess) {
throw std::runtime_error("Extract kernel launch failed");
}
}
"""
cpp_source = """
#include <torch/extension.h>
void fp4_gemv_cutlass(
torch::Tensor A,
torch::Tensor B,
torch::Tensor SFA,
torch::Tensor SFB,
torch::Tensor C_out,
int m, int k, int l,
int sfa_s0, int sfa_s1, int sfa_s2, int sfa_s3, int sfa_s4, int sfa_s5,
int sfb_s0, int sfb_s1, int sfb_s2, int sfb_s3, int sfb_s4, int sfb_s5
);
"""
_cuda_module = None
def get_cuda_module():
"""Compile and cache the CUDA module."""
global _cuda_module
if _cuda_module is None:
import os
# Get CUTLASS include path
cutlass_include = os.environ.get('CUTLASS_PATH', '/tmp/cutlass')
cutlass_include_dir = f"{cutlass_include}/include"
_cuda_module = load_inline(
name='fp4_gemv_cutlass_opt',
cpp_sources=[cpp_source],
cuda_sources=[cuda_source],
functions=['fp4_gemv_cutlass'],
extra_cuda_cflags=[
'-O3',
'--use_fast_math',
'-gencode=arch=compute_100a,code=sm_100a',
'-std=c++17',
f'-I{cutlass_include_dir}',
'-DCUTLASS_ARCH_MMA_SM100_SUPPORTED',
'-DCUTLASS_ARCH_MMA_SM100_ENABLED',
'-DCUTLASS_ARCH_MMA_SM100A_ENABLED',
'-DCUTE_ARCH_TMA_SM90_ENABLED',
'-DCUTE_ARCH_TCGEN05_TMEM_ENABLED',
'-DCUTLASS_ENABLE_TENSOR_CORE_MMA=1',
'-DCUTLASS_DEBUG_TRACE_LEVEL=0',
'-DCUTLASS_ENABLE_DEBUG_SYNCLOG=0',
'--expt-extended-lambda',
'--expt-relaxed-constexpr',
'-Xcudafe', '--diag_suppress=177'
],
verbose=True
)
return _cuda_module
def custom_kernel(data: input_t) -> output_t:
a, b, _, _, sfa_permuted, sfb_permuted, c = data
m, k_half, l = a.shape
k = k_half * 2
n_problem = 32 # Must match CUDA code
cuda_module = get_cuda_module()
# The scale factor strides are not actually used by CUTLASS (it computes its own layout)
# but we pass them for interface compatibility
sfa_strides = [0, 0, 0, 0, 0, 0]
sfb_strides = [0, 0, 0, 0, 0, 0]
# C_out will be filled directly by the extraction kernel in C++ layer
# Ensure c has the right shape for output
if c.dim() == 3 and c.size(1) == 1:
# Reshape to (m, l) for kernel, then restore view
c_reshaped = c.view(m, l)
cuda_module.fp4_gemv_cutlass(
a, b,
sfa_permuted, sfb_permuted,
c_reshaped,
m, k, l,
*sfa_strides,
*sfb_strides
)
else:
# Direct write to c (already (m, l) shape)
cuda_module.fp4_gemv_cutlass(
a, b,
sfa_permuted, sfb_permuted,
c,
m, k, l,
*sfa_strides,
*sfb_strides
)
return c
scrolls · 368 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