Skip to content
KernelIndex
Search⌘K

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
NVFP4 GEMMsuite of 3 cases
NVIDIA B200
29.6µs
#209 of 369
2025-12-14

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.

clusterusing ClusterShape = Shape<
fp4using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
fused-epilogueusing EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
warp-specializationusing 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 c
scrolls · 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 lines
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>
⋯ 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 configuration
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutATag = cutlass::layout::RowMajor;
constexpr int AlignmentA = 32;
⋯ 13 unchanged lines
using 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 == 128
using 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 lines
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());
+ 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 lines
Gemm 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 lines
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",
+ 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 c
No newline at end of file
scrolls · 143 diff lines total

Best evidence level for this revision: reported

JSON