Skip to content
KernelIndex
Search⌘K

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
NVFP4 GEMMsuite of 3 cases
NVIDIA B200
11.9µs
#95 of 369
2025-12-09

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.

clusterusing ClusterShape = Shape<_1, _2, _1>;
fp4using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
fused-epilogueusing CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
warp-specializationcutlass::epilogue::NoSmemWarpSpecialized1Sm

Kernel 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