submission 155451
snowclipsed · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 245 lines, June 9 Researcher Reciprocity License v1.0.
gemm_tma.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-155451?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:533e688f614f167d1ac4711479c6eaa26b385c4432923da9ae521c950065fdcd
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<fp4
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;fused-epilogue
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;warp-specialization
using KernelSchedule = cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100;Kernel source
gemm_tma.py245 lines
import torch
from torch.utils.cpp_extension import load_inline
import os
input_t = tuple
output_t = torch.Tensor
# Rethink for occupancy:
# - Use smaller N-tile (128x128x256) to create more CTAs and better fill SMs, especially for big N.
# - Keep 1x1x1 cluster to avoid reducing CTA count per wave.
# - Let the SM100 scheduler pick rasterization and swizzle heuristics.
cuda_source = r'''
#include <torch/extension.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include "cutlass/cutlass.h"
#include "cute/tensor.hpp"
#include "cutlass/tensor_ref.h"
#include "cutlass/epilogue/thread/linear_combination.h"
#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/gemm/kernel/tile_scheduler_params.h"
#include "cutlass/util/packed_stride.hpp"
using namespace cute;
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
// NVFP4 configuration
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 ElementD = cutlass::half_t;
using ElementC = cutlass::half_t;
using LayoutCTag = cutlass::layout::RowMajor;
using LayoutDTag = cutlass::layout::RowMajor;
constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
using ElementAccumulator = float;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
// For occupancy: choose 128x128x256 to increase CTA count (vs 256 N-tile).
#ifndef MMA_N_TILE
#define MMA_N_TILE 128
#endif
#if !defined(CLUSTER_M)
#define CLUSTER_M 1
#endif
#if !defined(CLUSTER_N)
#define CLUSTER_N 1
#endif
template<int M> struct ClusterM;
template<> struct ClusterM<1> { using type = _1; };
template<> struct ClusterM<2> { using type = _2; };
template<> struct ClusterM<4> { using type = _4; };
template<> struct ClusterM<8> { using type = _8; };
template<int N> struct ClusterN;
template<> struct ClusterN<1> { using type = _1; };
template<> struct ClusterN<2> { using type = _2; };
template<> struct ClusterN<4> { using type = _4; };
template<> struct ClusterN<8> { using type = _8; };
#if MMA_N_TILE == 128
using MmaTileShape = Shape<_128,_128,_256>;
#elif MMA_N_TILE == 192
using MmaTileShape = Shape<_128,_192,_256>;
#elif MMA_N_TILE == 256
using MmaTileShape = Shape<_128,_256,_256>;
#else
#error "Unsupported MMA_N_TILE. Use 128, 192, or 256."
#endif
using ClusterShape = Shape<
typename ClusterM<CLUSTER_M>::type,
typename ClusterN<CLUSTER_N>::type,
_1
>;
// 1SM NVF4 TN schedules
using KernelSchedule = cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100;
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag, OperatorClass,
MmaTileShape, ClusterShape,
cutlass::epilogue::collective::EpilogueTileAuto,
ElementAccumulator, ElementAccumulator,
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 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;
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});
StrideC stride_C = cutlass::make_cute_packed_stride(StrideC{}, {m, n, 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));
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* c_ptr = reinterpret_cast<ElementD*>(c.data_ptr());
typename Gemm::Arguments arguments {
cutlass::gemm::GemmUniversalMode::kGemm,
{m, n, k, l},
{a_ptr, stride_A, b_ptr, stride_B, sfa_ptr, layout_SFA, sfb_ptr, layout_SFB},
{{1.0f, 0.0f}, nullptr, stride_C, c_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()));
auto status = gemm.can_implement(arguments);
TORCH_CHECK(status == cutlass::Status::kSuccess, "CUTLASS cannot implement this GEMM");
status = gemm.initialize(arguments, workspace.data_ptr());
TORCH_CHECK(status == cutlass::Status::kSuccess, "CUTLASS initialization failed");
status = gemm.run();
TORCH_CHECK(status == cutlass::Status::kSuccess, "CUTLASS kernel failed");
cudaDeviceSynchronize();
}
#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
) {
TORCH_CHECK(false, "SM100 not supported in this build");
}
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("nvfp4_gemm", &nvfp4_gemm, "NVFP4 Block-Scaled 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")
# Build with smaller N tile (128) and 1x1 cluster to increase CTA count and occupancy.
module = load_inline(
name="nvfp4_gemm_module_n128",
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",
"-DMMA_N_TILE=128",
"-DCLUSTER_M=1",
"-DCLUSTER_N=1",
"-O3",
# Optional: try to improve L2 caching behavior; may help for large N
"-Xptxas=-dlcm=ca",
],
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) # FP4 packed K/2 -> K
module.nvfp4_gemm(a, b, sfa_perm, sfb_perm, c, m, n, k, l)
return cscrolls · 245 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 155383.
⋯ 4 unchanged linesinput_t = tupleoutput_t = torch.Tensor+ # Rethink for occupancy:+ # - Use smaller N-tile (128x128x256) to create more CTAs and better fill SMs, especially for big N.+ # - Keep 1x1x1 cluster to avoid reducing CTA count per wave.+ # - Let the SM100 scheduler pick rasterization and swizzle heuristics.cuda_source = r'''#include <torch/extension.h>#include <cuda_runtime.h>⋯ 16 unchanged lines#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)- // NVFP4 configuration:- // - nv_float4_t<float_e2m1_t> uses float_ue4m3_t scale factors (unsigned E4M3)- // - Vector size is 16 (one scale factor per 16 FP4 elements)- // - Layout is TN (A=RowMajor, B=ColumnMajor)+ // NVFP4 configurationusing ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;using LayoutATag = cutlass::layout::RowMajor;constexpr int AlignmentA = 32;⋯ 13 unchanged linesusing ArchTag = cutlass::arch::Sm100;using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;- // Tile shape for NVF4 - use 128x128x256 with 1x1x1 cluster for single SM+ // For occupancy: choose 128x128x256 to increase CTA count (vs 256 N-tile).+ #ifndef MMA_N_TILE+ #define MMA_N_TILE 128+ #endif++ #if !defined(CLUSTER_M)+ #define CLUSTER_M 1+ #endif+ #if !defined(CLUSTER_N)+ #define CLUSTER_N 1+ #endif++ template<int M> struct ClusterM;+ template<> struct ClusterM<1> { using type = _1; };+ template<> struct ClusterM<2> { using type = _2; };+ template<> struct ClusterM<4> { using type = _4; };+ template<> struct ClusterM<8> { using type = _8; };++ template<int N> struct ClusterN;+ template<> struct ClusterN<1> { using type = _1; };+ template<> struct ClusterN<2> { using type = _2; };+ template<> struct ClusterN<4> { using type = _4; };+ template<> struct ClusterN<8> { using type = _8; };++ #if MMA_N_TILE == 128using MmaTileShape = Shape<_128,_128,_256>;- using ClusterShape = Shape<_1,_1,_1>;+ #elif MMA_N_TILE == 192+ using MmaTileShape = Shape<_128,_192,_256>;+ #elif MMA_N_TILE == 256+ using MmaTileShape = Shape<_128,_256,_256>;+ #else+ #error "Unsupported MMA_N_TILE. Use 128, 192, or 256."+ #endif- // NVF4-specific kernel and epilogue schedules- using KernelSchedule = cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100;+ using ClusterShape = Shape<+ typename ClusterM<CLUSTER_M>::type,+ typename ClusterN<CLUSTER_N>::type,+ _1+ >;++ // 1SM NVF4 TN schedules+ using KernelSchedule = cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100;using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<⋯ 50 unchanged linesLayoutSFB layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(m, n, k, 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* c_ptr = reinterpret_cast<ElementD*>(c.data_ptr());+ 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* c_ptr = reinterpret_cast<ElementD*>(c.data_ptr());typename Gemm::Arguments arguments {cutlass::gemm::GemmUniversalMode::kGemm,⋯ 4 unchanged linesGemm gemm;size_t workspace_size = Gemm::get_workspace_size(arguments);-- auto workspace = torch::empty({static_cast<long>(workspace_size)},++ auto workspace = torch::empty({static_cast<long>(workspace_size)},torch::TensorOptions().dtype(torch::kUInt8).device(a.device()));auto status = gemm.can_implement(arguments);⋯ 37 unchanged linescutlass_path = os.environ.get("CUTLASS_PATH", "/usr/local/cutlass")cuda_include = os.environ.get("CUDA_INCLUDE_DIR", "/usr/local/cuda/include")+ # Build with smaller N tile (128) and 1x1 cluster to increase CTA count and occupancy.module = load_inline(- name="nvfp4_gemm_module",+ name="nvfp4_gemm_module_n128",cpp_sources=[cpp_source],cuda_sources=[cuda_source],extra_include_paths=[⋯ 5 unchanged lines"-std=c++17","-arch=sm_100a","-DCUTLASS_ARCH_MMA_SM100_SUPPORTED=1",+ "-DMMA_N_TILE=128",+ "-DCLUSTER_M=1",+ "-DCLUSTER_N=1","-O3",+ # Optional: try to improve L2 caching behavior; may help for large N+ "-Xptxas=-dlcm=ca",],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) # k/2 due to FP4 packing-+ k = int(a.shape[1] * 2) # FP4 packed K/2 -> K+module.nvfp4_gemm(a, b, sfa_perm, sfb_perm, c, m, n, k, l)return cNo newline at end of file
scrolls · 143 diff lines total
Best evidence level for this revision: reported
JSON