Skip to content
KernelIndex
Search⌘K

submission 156306

snowclipsed · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 188 lines, June 9 Researcher Reciprocity License v1.0.

gemm_tma.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-156306?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
14.5µs
#155 of 369
2025-12-14

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:54f1c355b89a7df423c034e4096650d105ecb3ca71075ad8a43548ae5764d80c
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<_1, _1, _1>;
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.py188 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/tensor_ref.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/util/packed_stride.hpp"

using namespace cute;

#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;

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 = 8;
constexpr int AlignmentC = 8;

using ElementAccumulator = float;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;

// Smallest supported tile for NVF4: 128x128x256
using MmaTileShape = Shape<_128, _128, _256>;
using ClusterShape = Shape<_1, _1, _1>;

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: ", cutlass::cutlassGetStatusString(status));

    status = gemm.initialize(arguments, workspace.data_ptr());
    TORCH_CHECK(status == cutlass::Status::kSuccess, 
        "CUTLASS init failed: ", cutlass::cutlassGetStatusString(status));

    status = gemm.run();
    TORCH_CHECK(status == cutlass::Status::kSuccess, 
        "CUTLASS run failed: ", cutlass::cutlassGetStatusString(status));
}

#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");
}

#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",
    ],
    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)
    module.nvfp4_gemm(a, b, sfa_perm, sfb_perm, c, m, n, k, l)
    return c
scrolls · 188 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 155451.

⋯ 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>
- #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;
+ // 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;
- using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
- using LayoutBTag = cutlass::layout::ColumnMajor;
- constexpr int AlignmentB = 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 ElementD = cutlass::half_t;
+ using ElementC = cutlass::half_t;
+ using LayoutCTag = cutlass::layout::RowMajor;
+ using LayoutDTag = cutlass::layout::RowMajor;
+ constexpr int AlignmentD = 8;
+ constexpr int AlignmentC = 8;
- using ElementAccumulator = float;
- using ArchTag = cutlass::arch::Sm100;
- using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
+ 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
+ // Smallest supported tile for NVF4: 128x128x256
+ using MmaTileShape = Shape<_128, _128, _256>;
+ using ClusterShape = Shape<_1, _1, _1>;
- #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 KernelSchedule = cutlass::gemm::KernelTmaWarpSpecialized1SmNvf4Sm100;
using EpilogueSchedule = cutlass::epilogue::TmaWarpSpecialized1SmNvf4;
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
⋯ 12 unchanged lines
ElementB, LayoutBTag, AlignmentB,
ElementAccumulator,
MmaTileShape, ClusterShape,
- cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
+ cutlass::gemm::collective::StageCountAutoCarveout<
+ static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
KernelSchedule
>::CollectiveOp;
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
- Shape<int,int,int,int>,
+ 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 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
- ) {
+ 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});
⋯ 4 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 {
+ 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},
⋯ 2 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);
- TORCH_CHECK(status == cutlass::Status::kSuccess, "CUTLASS cannot implement this GEMM");
+ TORCH_CHECK(status == cutlass::Status::kSuccess,
+ "CUTLASS cannot implement: ", cutlass::cutlassGetStatusString(status));
status = gemm.initialize(arguments, workspace.data_ptr());
- TORCH_CHECK(status == cutlass::Status::kSuccess, "CUTLASS initialization failed");
+ TORCH_CHECK(status == cutlass::Status::kSuccess,
+ "CUTLASS init failed: ", cutlass::cutlassGetStatusString(status));
status = gemm.run();
- TORCH_CHECK(status == cutlass::Status::kSuccess, "CUTLASS kernel failed");
-
- cudaDeviceSynchronize();
+ TORCH_CHECK(status == cutlass::Status::kSuccess,
+ "CUTLASS run failed: ", cutlass::cutlassGetStatusString(status));
}
#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");
+ 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");
}
#endif
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
- m.def("nvfp4_gemm", &nvfp4_gemm, "NVFP4 Block-Scaled GEMM");
+ 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
- );
+ 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",
+ name="nvfp4_gemm_module",
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) # FP4 packed K/2 -> K
-
+ k = int(a.shape[1] * 2)
module.nvfp4_gemm(a, b, sfa_perm, sfb_perm, c, m, n, k, l)
return c
No newline at end of file
scrolls · 279 diff lines total

Best evidence level for this revision: reported

JSON