submission 158766
snowclipsed · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 165 lines, June 9 Researcher Reciprocity License v1.0.
gemm_tma.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-158766?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:dd1d66fed39775d3861ab42d9d866769247bdf4686e01db25e1b75582813cf4f
license declaredunknown
license concludedunknown
authorssnowclipsed
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 ElementCompute = cutlass::half_t; // Use half for epilogue computewarp-specialization
using KernelSchedule = cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100;Kernel source
gemm_tma.py165 lines
import torch
from torch.utils.cpp_extension import load_inline
import os
input_t = tuple
output_t = torch.Tensor
cuda_source = r'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include "cutlass/cutlass.h"
#include "cute/tensor.hpp"
#include "cutlass/gemm/dispatch_policy.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"
using namespace cute;
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutATag = cutlass::layout::RowMajor;
constexpr int AlignmentA = 64;
using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutBTag = cutlass::layout::ColumnMajor;
constexpr int AlignmentB = 64;
using ElementD = cutlass::half_t;
using ElementC = void;
using LayoutCTag = cutlass::layout::RowMajor;
using LayoutDTag = cutlass::layout::RowMajor;
constexpr int AlignmentD = 16;
constexpr int AlignmentC = 1;
using ElementAccumulator = float;
using ElementCompute = cutlass::half_t; // Use half for epilogue compute
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
using MmaTileShape = Shape<_128, _64, _256>;
using ClusterShape = Shape<_1, _2, _1>;
using KernelSchedule = cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100;
using EpilogueSchedule = cutlass::epilogue::NoSmemWarpSpecialized1Sm;
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag, OperatorClass,
MmaTileShape, ClusterShape,
cutlass::epilogue::collective::EpilogueTileAuto,
ElementAccumulator, ElementCompute,
ElementC, LayoutCTag, AlignmentC,
ElementD, LayoutDTag, AlignmentD,
EpilogueSchedule
>::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))>,
KernelSchedule
>::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 StrideB = typename Gemm::GemmKernel::StrideB;
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;
void nvfp4_gemm(
torch::Tensor a, torch::Tensor b,
torch::Tensor sfa_perm, torch::Tensor sfb_perm,
torch::Tensor c, int m, int n, int k, int l)
{
StrideA stride_A = cutlass::make_cute_packed_stride(StrideA{}, {m, k, l});
StrideB stride_B = cutlass::make_cute_packed_stride(StrideB{}, {n, k, l});
StrideD stride_D = cutlass::make_cute_packed_stride(StrideD{}, {m, n, l});
LayoutSFA layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(m, n, k, l));
LayoutSFB layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(m, n, k, l));
typename Gemm::Arguments arguments{
cutlass::gemm::GemmUniversalMode::kGemm,
{m, n, k, l},
{reinterpret_cast<typename ElementA::DataType*>(a.data_ptr()), stride_A,
reinterpret_cast<typename ElementB::DataType*>(b.data_ptr()), stride_B,
reinterpret_cast<typename ElementA::ScaleFactorType*>(sfa_perm.data_ptr()), layout_SFA,
reinterpret_cast<typename ElementB::ScaleFactorType*>(sfb_perm.data_ptr()), layout_SFB},
{{cutlass::half_t(1.0f), cutlass::half_t(0.0f)}, nullptr, {},
reinterpret_cast<ElementD*>(c.data_ptr()), stride_D}
};
Gemm gemm;
size_t workspace_size = Gemm::get_workspace_size(arguments);
auto workspace = torch::empty({static_cast<long>(workspace_size)},
torch::TensorOptions().dtype(torch::kUInt8).device(a.device()));
gemm.initialize(arguments, workspace.data_ptr());
gemm.run();
}
#else
void nvfp4_gemm(torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor,
torch::Tensor, int, int, int, int) {
TORCH_CHECK(false, "SM100 not supported");
}
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("nvfp4_gemm", &nvfp4_gemm, "NVFP4 GEMM");
}
'''
cpp_source = """
#include <torch/extension.h>
void nvfp4_gemm(torch::Tensor a, torch::Tensor b,
torch::Tensor sfa_perm, torch::Tensor sfb_perm,
torch::Tensor c, int m, int n, int k, int l);
"""
cutlass_path = os.environ.get("CUTLASS_PATH", "/usr/local/cutlass")
cuda_include = os.environ.get("CUDA_INCLUDE_DIR", "/usr/local/cuda/include")
module = load_inline(
name="nvfp4_gemm_module",
cpp_sources=[cpp_source],
cuda_sources=[cuda_source],
extra_include_paths=[
f"{cutlass_path}/include",
f"{cutlass_path}/tools/util/include",
cuda_include,
],
extra_cuda_cflags=[
"-std=c++17",
"-arch=sm_100a",
"-DCUTLASS_ARCH_MMA_SM100_SUPPORTED=1",
"-O3",
"--use_fast_math",
],
verbose=False,
)
def custom_kernel(data: input_t) -> output_t:
a, b, sfa, sfb, sfa_perm, sfb_perm, c = data
m, n, l = c.shape[0], c.shape[1], c.shape[2]
k = a.shape[1] * 2
module.nvfp4_gemm(a, b, sfa_perm, sfb_perm, c, m, n, k, l)
return cscrolls · 165 lines total
Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0
Changes from previous submission
Against this author's previous submission submission 158340.
⋯ 9 unchanged lines#include <cuda_runtime.h>#include "cutlass/cutlass.h"#include "cute/tensor.hpp"- #include "cutlass/tensor_ref.h"#include "cutlass/gemm/dispatch_policy.hpp"#include "cutlass/gemm/collective/collective_builder.hpp"#include "cutlass/epilogue/collective/collective_builder.hpp"⋯ 6 unchanged lines#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)- // NVFP4 with float_ue4m3_t scale factors (VS=16, block16)using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;using LayoutATag = cutlass::layout::RowMajor;- constexpr int AlignmentA = 32;+ constexpr int AlignmentA = 64;using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;using LayoutBTag = cutlass::layout::ColumnMajor;- constexpr int AlignmentB = 32;+ constexpr int AlignmentB = 64;using ElementD = cutlass::half_t;- using ElementC = void; // No bias matrix - pure D = A*B+ using ElementC = void;using LayoutCTag = cutlass::layout::RowMajor;using LayoutDTag = cutlass::layout::RowMajor;constexpr int AlignmentD = 16;- constexpr int AlignmentC = 1; // void type+ constexpr int AlignmentC = 1;using ElementAccumulator = float;+ using ElementCompute = cutlass::half_t; // Use half for epilogue compute+using ArchTag = cutlass::arch::Sm100;using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;using MmaTileShape = Shape<_128, _64, _256>;using ClusterShape = Shape<_1, _2, _1>;- using KernelSchedule = cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100;+ using KernelSchedule = cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100;using EpilogueSchedule = cutlass::epilogue::NoSmemWarpSpecialized1Sm;using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<ArchTag, OperatorClass,MmaTileShape, ClusterShape,cutlass::epilogue::collective::EpilogueTileAuto,- ElementAccumulator, ElementAccumulator,+ ElementAccumulator, ElementCompute,ElementC, LayoutCTag, AlignmentC,ElementD, LayoutDTag, AlignmentD,EpilogueSchedule⋯ 17 unchanged linesvoid>;using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;-using StrideA = typename Gemm::GemmKernel::StrideA;using StrideB = typename Gemm::GemmKernel::StrideB;using StrideD = typename Gemm::GemmKernel::StrideD;⋯ 1 unchanged linesusing LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFB;using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;- // Cached state for repeated calls- static thread_local Gemm* cached_gemm = nullptr;- static thread_local void* cached_workspace = nullptr;- static thread_local int cached_m = 0, cached_n = 0, cached_k = 0, cached_l = 0;- static thread_local StrideA cached_stride_A;- static thread_local StrideB cached_stride_B;- static thread_local StrideD cached_stride_D;- static thread_local LayoutSFA cached_layout_SFA;- static thread_local LayoutSFB cached_layout_SFB;-void nvfp4_gemm(torch::Tensor a, torch::Tensor b,torch::Tensor sfa_perm, torch::Tensor sfb_perm,torch::Tensor c, int m, int n, int k, int l){- bool dims_changed = (m != cached_m || n != cached_n || k != cached_k || l != cached_l);-- if (dims_changed || cached_gemm == nullptr) {- cached_stride_A = cutlass::make_cute_packed_stride(StrideA{}, {m, k, l});- cached_stride_B = cutlass::make_cute_packed_stride(StrideB{}, {n, k, l});- cached_stride_D = cutlass::make_cute_packed_stride(StrideD{}, {m, n, l});- cached_layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(m, n, k, l));- cached_layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(m, n, k, l));- }+ StrideA stride_A = cutlass::make_cute_packed_stride(StrideA{}, {m, k, l});+ StrideB stride_B = cutlass::make_cute_packed_stride(StrideB{}, {n, k, l});+ StrideD stride_D = cutlass::make_cute_packed_stride(StrideD{}, {m, n, l});- auto* a_ptr = reinterpret_cast<typename ElementA::DataType*>(a.data_ptr());- auto* b_ptr = reinterpret_cast<typename ElementB::DataType*>(b.data_ptr());- auto* sfa_ptr = reinterpret_cast<typename ElementA::ScaleFactorType*>(sfa_perm.data_ptr());- auto* sfb_ptr = reinterpret_cast<typename ElementB::ScaleFactorType*>(sfb_perm.data_ptr());- auto* d_ptr = reinterpret_cast<ElementD*>(c.data_ptr());+ LayoutSFA layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(m, n, k, l));+ LayoutSFB layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(m, n, k, l));typename Gemm::Arguments arguments{cutlass::gemm::GemmUniversalMode::kGemm,{m, n, k, l},- {a_ptr, cached_stride_A, b_ptr, cached_stride_B, sfa_ptr, cached_layout_SFA, sfb_ptr, cached_layout_SFB},- {{1.0f, 0.0f}, nullptr, {}, d_ptr, cached_stride_D}+ {reinterpret_cast<typename ElementA::DataType*>(a.data_ptr()), stride_A,+ reinterpret_cast<typename ElementB::DataType*>(b.data_ptr()), stride_B,+ reinterpret_cast<typename ElementA::ScaleFactorType*>(sfa_perm.data_ptr()), layout_SFA,+ reinterpret_cast<typename ElementB::ScaleFactorType*>(sfb_perm.data_ptr()), layout_SFB},+ {{cutlass::half_t(1.0f), cutlass::half_t(0.0f)}, nullptr, {},+ reinterpret_cast<ElementD*>(c.data_ptr()), stride_D}};-- if (dims_changed || cached_gemm == nullptr) {- delete cached_gemm;- cached_gemm = new Gemm();-- size_t workspace_size = Gemm::get_workspace_size(arguments);- if (workspace_size > 0) {- cudaMalloc(&cached_workspace, workspace_size);- } else {- cached_workspace = nullptr;- }-- auto status = cached_gemm->initialize(arguments, cached_workspace);- TORCH_CHECK(status == cutlass::Status::kSuccess,- "CUTLASS init failed: ", cutlass::cutlassGetStatusString(status));-- cached_m = m; cached_n = n; cached_k = k; cached_l = l;- } else {- auto status = cached_gemm->update(arguments, cached_workspace);- TORCH_CHECK(status == cutlass::Status::kSuccess,- "CUTLASS update failed: ", cutlass::cutlassGetStatusString(status));- }- auto status = cached_gemm->run();- TORCH_CHECK(status == cutlass::Status::kSuccess,- "CUTLASS run failed: ", cutlass::cutlassGetStatusString(status));+ Gemm gemm;+ size_t workspace_size = Gemm::get_workspace_size(arguments);+ auto workspace = torch::empty({static_cast<long>(workspace_size)},+ torch::TensorOptions().dtype(torch::kUInt8).device(a.device()));++ gemm.initialize(arguments, workspace.data_ptr());+ gemm.run();}#else- void nvfp4_gemm(torch::Tensor a, torch::Tensor b,- torch::Tensor sfa_perm, torch::Tensor sfb_perm,- torch::Tensor c, int m, int n, int k, int l) {+ void nvfp4_gemm(torch::Tensor, torch::Tensor, torch::Tensor, torch::Tensor,+ torch::Tensor, int, int, int, int) {TORCH_CHECK(false, "SM100 not supported");}#endif⋯ 27 unchanged lines"-arch=sm_100a","-DCUTLASS_ARCH_MMA_SM100_SUPPORTED=1","-O3",+ "--use_fast_math",],verbose=False,)def custom_kernel(data: input_t) -> output_t:a, b, sfa, sfb, sfa_perm, sfb_perm, c = data- m = int(c.shape[0])- n = int(c.shape[1])- l = int(c.shape[2])- k = int(a.shape[1] * 2)+ m, n, l = c.shape[0], c.shape[1], c.shape[2]+ k = a.shape[1] * 2module.nvfp4_gemm(a, b, sfa_perm, sfb_perm, c, m, n, k, l)return cNo newline at end of file
scrolls · 178 diff lines total
Best evidence level for this revision: reported
JSON