submission 137115
mdouglas · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 206 lines, June 9 Researcher Reciprocity License v1.0.
gemm3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-137115?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:260929b1e535dab71df66b04a0192f7bb7fb9fdd01d6cd7951a2f68d2ecc9b25
license declaredunknown
license concludedunknown
authorsmdouglas
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
using ClusterShape = Shape<_1, _2, _1>;fp4
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;fused-epilogue
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<warp-specialization
cutlass::epilogue::NoSmemWarpSpecialized1SmKernel source
gemm3.py206 lines
from torch.utils.cpp_extension import load_inline
cuda_source = """
#include <torch/extension.h>
#include <cuda_runtime.h>
#include "cutlass/cutlass.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/dispatch_policy.hpp"
#include "cutlass/util/packed_stride.hpp"
using namespace cute;
// Kernel configuration for NVFP4 block-scaled GEMM with FP16 output
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 LayoutCTag = cutlass::layout::RowMajor;
constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
using ElementD = ElementC;
using LayoutDTag = LayoutCTag;
constexpr int AlignmentD = AlignmentC;
using ElementSFA = cutlass::float_ue4m3_t;
using ElementSFB = cutlass::float_ue4m3_t;
using ElementAccumulator = float;
using ElementCompute = float;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
using MmaTileShape = Shape<_128, _64, _256>;
using ClusterShape = Shape<_1, _2, _1>;
using PerSmTileShapeMNK = Shape<_128, _64, _256>;
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag,
OperatorClass,
PerSmTileShapeMNK,
ClusterShape,
cutlass::epilogue::collective::EpilogueTileAuto,
ElementAccumulator,
ElementC, //ElementCompute,
ElementC,
LayoutCTag,
AlignmentC,
ElementD,
LayoutDTag,
AlignmentD,
cutlass::epilogue::NoSmemWarpSpecialized1Sm
>::CollectiveOp;
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
ArchTag,
OperatorClass,
ElementA,
LayoutATag,
AlignmentA,
ElementB,
LayoutBTag,
AlignmentB,
ElementAccumulator, //float32
MmaTileShape,
ClusterShape,
cutlass::gemm::collective::StageCount<8>,
cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100
>::CollectiveOp;
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
Shape<int, int, int, int>,
CollectiveMainloop,
CollectiveEpilogue,
void
>;
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
using StrideA = typename Gemm::GemmKernel::StrideA;
using LayoutA = decltype(cute::make_layout(make_shape(0, 0, 0), StrideA{}));
using LayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFA;
using StrideB = typename Gemm::GemmKernel::StrideB;
using LayoutB = decltype(cute::make_layout(make_shape(0, 0, 0), StrideB{}));
using LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFB;
using StrideC = typename Gemm::GemmKernel::StrideC;
using LayoutC = decltype(cute::make_layout(make_shape(0, 0, 0), StrideC{}));
using StrideD = typename Gemm::GemmKernel::StrideD;
using LayoutD = decltype(cute::make_layout(make_shape(0, 0, 0), StrideD{}));
torch::Tensor run_gemm(std::vector<torch::Tensor> data) {
using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;;
auto a = data[1]; // [M, K/2] uint8 packed FP4
auto b = data[0]; // [N, K/2] uint8 packed FP4
auto scale_a = data[5]; // Permuted FP8 scales
auto scale_b = data[4]; // Permuted FP8 scales
//auto c = data[6]; // [M, N] FP16 output
const int m = a.size(0);
const int k = a.size(1) * 2; // Unpacked K dimension
const int n = b.size(0);
auto c = torch::empty({m, n, 1}, torch::dtype(torch::kFloat16).device(a.device()));
auto a_ptr = static_cast<typename Gemm::ElementA const*>(a.data_ptr());
auto b_ptr = static_cast<typename Gemm::ElementB const*>(b.data_ptr());
auto c_ptr = static_cast<ElementD*>(c.data_ptr());
auto sfa_ptr = static_cast<ElementSFA const*>(scale_a.data_ptr());
auto sfb_ptr = static_cast<ElementSFB const*>(scale_b.data_ptr());
auto stride_A = cutlass::make_cute_packed_stride(StrideA{}, {m, k, 1});
auto stride_B = cutlass::make_cute_packed_stride(StrideB{}, {n, k, 1});
auto stride_D = cutlass::make_cute_packed_stride(StrideD{}, {m, n, 1});
auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(
cute::make_shape(m, n, k, 1)
);
auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(
cute::make_shape(m, n, k, 1)
);
typename Gemm::Arguments arguments{
cutlass::gemm::GemmUniversalMode::kGemm,
{m, n, k, 1},
{ // Mainloop
a_ptr, stride_A,
b_ptr, stride_B,
sfa_ptr, layout_SFA,
sfb_ptr, layout_SFB
},
{ // Epilogue
{}, // epilogue.thread
c_ptr, stride_D,
c_ptr, stride_D
}
};
Gemm gemm_op;
// Allocate workspace if needed
void* workspace_ptr = nullptr;
size_t workspace_size = gemm_op.get_workspace_size(arguments);
torch::Tensor workspace;
if (workspace_size > 0) {
workspace = torch::empty({static_cast<int64_t>(workspace_size)},
torch::dtype(torch::kUInt8).device(a.device()));
workspace_ptr = workspace.data_ptr();
}
gemm_op.initialize(arguments, workspace_ptr);
gemm_op.run();
return c.transpose(0, 1);
}
"""
cpp_source = """
torch::Tensor run_gemm(std::vector<torch::Tensor> data);
"""
module = load_inline(
name='nvfp4_gemm',
cpp_sources=cpp_source,
cuda_sources=cuda_source,
functions=['run_gemm'],
with_cuda=True,
extra_cflags=[
'-O3',
'-std=c++20',
'-march=native',
],
extra_cuda_cflags=[
'-O3',
'--use_fast_math',
'--extra-device-vectorization',
'--maxrregcount=128',
'--restrict',
'-arch=sm_100a',
'-Xptxas=-v',
'-lineinfo',
'-std=c++20',
'-U__CUDA_NO_HALF_OPERATORS__',
'-U__CUDA_NO_HALF_CONVERSIONS__',
],
verbose=True,
)
custom_kernel = module.run_gemm
scrolls · 206 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