Skip to content
KernelIndex
Search⌘K

submission 180670

whoknows · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

sol_prob2.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-180670?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.4µs
#154 of 369
2025-12-19

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:cb7de0646e6eef925249fa63bcfba50848acb8fc1129d89f91dd9f55593aec18
license declaredunknown
license concludedunknown
authorswhoknows
imported2026-08-26

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<

Kernel source

sol_prob2.py318 lines
import torch
from torch.utils.cpp_extension import load_inline
from typing import TypeVar
import os
from pathlib import Path

# Type definitions matching the task
input_t = TypeVar("input_t", bound=tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor])
output_t = TypeVar("output_t", bound=torch.Tensor)

# Get CUTLASS include path - check environment variable first, then common locations
# CUTLASS_DIR = os.environ.get("CUTLASS_DIR")
# if not CUTLASS_DIR:
#     # Try common locations
#     possible_paths = [
#         str(Path.home() / "cutlass"),   # ~/cutlass
#         "/usr/local/cutlass",            # System install
#         "/opt/cutlass",                  # Alternative system install
#         str(Path(__file__).parent.parent / "cutlass"),  # ../cutlass from script location
#     ]
#     for path in possible_paths:
#         if Path(path).exists() and (Path(path) / "include" / "cutlass").exists():
#             CUTLASS_DIR = path
#             break
    
#     if not CUTLASS_DIR:
#         raise RuntimeError(
#             "CUTLASS not found. Please set CUTLASS_DIR environment variable or install CUTLASS in one of:\n" +
#             "\n".join(f"  - {p}" for p in possible_paths)
#         )

# ---- C++ stub: declare the function so load_inline can bind it ----
gemm_cpp = r"""
#include <torch/extension.h>

// Forward declaration so PyTorch can bind it
// Accept raw byte tensors to avoid PyTorch type checking
torch::Tensor cuda_nvfp4_gemm(
    torch::Tensor A,        // underlying storage as bytes
    torch::Tensor B,        // underlying storage as bytes
    torch::Tensor SFA,      // float8_e4m3fn
    torch::Tensor SFB,      // float8_e4m3fn
    torch::Tensor C         // float16
);
"""

# ---- CUDA source: CUTLASS GEMM wrapper ----
gemm_cuda = r"""
#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.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"

#include "cutlass/gemm/collective/sm100_blockscaled_mma_mixed_tma_cpasync_warpspecialized.hpp"

using namespace cute;

// GEMM kernel configurations from 72a_blackwell_nvfp4_bf16_gemm.cu
// Modified to output FP16 instead of BF16
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;

// using MmaTileShape        = Shape<_256,_256,_256>;
// using MmaTileShape        = Shape<_128,_128,_256>;
using MmaTileShape        = Shape<_128,_128,_128>;
// using ClusterShape        = Shape<_2,_4,_1>;
// using ClusterShape        = Shape<_1,_1,_1>;
using ClusterShape        = Shape<_1,_2,_1>;

using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
    ArchTag, OperatorClass,
    MmaTileShape, ClusterShape,
    cutlass::epilogue::collective::EpilogueTileAuto,
    ElementAccumulator, ElementAccumulator,
    ElementC, LayoutCTag, AlignmentC,
    ElementD, LayoutDTag, AlignmentD,
    cutlass::epilogue::collective::EpilogueScheduleAuto
  >::CollectiveOp;

// using StageCount = cutlass::gemm::collective::StageCount<5>; // worked for <128,128,256> with <1,1,1>

using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
    ArchTag, OperatorClass,
    ElementA, LayoutATag, AlignmentA,
    ElementB, LayoutBTag, AlignmentB,
    ElementAccumulator,
    MmaTileShape, ClusterShape,
//    StageCount,
    cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
    cutlass::gemm::collective::KernelScheduleAuto
  >::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 cuda_nvfp4_gemm(
    torch::Tensor A,        // [m, k/2, l] - underlying storage as bytes
    torch::Tensor B,        // [n, k/2, l] - underlying storage as bytes
    torch::Tensor SFA,      // [32, 4, rest_m, 4, rest_k, l] in float8_e4m3fn
    torch::Tensor SFB,      // [32, 4, rest_n, 4, rest_k, l] in float8_e4m3fn
    torch::Tensor C)        // [m, n, l] in float16 (output - preallocated)
{
    using namespace cute;
    using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
    
    // Extract dimensions
    const int m = A.size(0);
    const int k_half = A.size(1);
    const int k = k_half * 2;  // Actual K dimension
    const int l = A.size(2);
    const int n = B.size(0);
    
    // Get underlying byte pointers
    // A and B are stored as packed bytes, we'll reinterpret them as FP4 data
    void* a_base = A.data_ptr();
    void* b_base = B.data_ptr();
    
    // Create layouts
    StrideA stride_A = cutlass::make_cute_packed_stride(StrideA{}, {m, k, 1});
    StrideB stride_B = cutlass::make_cute_packed_stride(StrideB{}, {n, k, 1});
    StrideC stride_C = cutlass::make_cute_packed_stride(StrideC{}, {m, n, 1});
    StrideD stride_D = cutlass::make_cute_packed_stride(StrideD{}, {m, n, 1});
    
    LayoutA layout_A = make_layout(make_shape(m, k, 1), stride_A);
    LayoutB layout_B = make_layout(make_shape(n, k, 1), stride_B);
    LayoutC layout_C = make_layout(make_shape(m, n, 1), stride_C);
    LayoutD layout_D = make_layout(make_shape(m, n, 1), stride_D);
    LayoutSFA layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(m, n, k, 1));
    LayoutSFB layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(m, n, k, 1));
    
    // Alpha and beta for computation
    float alpha = 1.0f;
    float beta = 0.0f;
    
    // Calculate strides in bytes for A and B
    size_t a_batch_stride_bytes = A.stride(2);
    size_t b_batch_stride_bytes = B.stride(2);
    
    // Process each batch element
    for (int batch = 0; batch < l; batch++) {
        // Get pointers for this batch - cast from void* to the FP4 data type
        auto* a_ptr = reinterpret_cast<typename ElementA::DataType*>(
            static_cast<char*>(a_base) + batch * a_batch_stride_bytes);
        auto* b_ptr = reinterpret_cast<typename ElementB::DataType*>(
            static_cast<char*>(b_base) + batch * b_batch_stride_bytes);
        auto* sfa_ptr = reinterpret_cast<typename ElementA::ScaleFactorType*>(
            SFA.data_ptr<at::Float8_e4m3fn>() + batch * SFA.stride(5));
        auto* sfb_ptr = reinterpret_cast<typename ElementB::ScaleFactorType*>(
            SFB.data_ptr<at::Float8_e4m3fn>() + batch * SFB.stride(5));
        auto* c_ptr = reinterpret_cast<cutlass::half_t*>(C.data_ptr<at::Half>() + batch * C.stride(2));
        auto* d_ptr = reinterpret_cast<cutlass::half_t*>(C.data_ptr<at::Half>() + batch * C.stride(2));
        
        // Create GEMM arguments
        typename Gemm::Arguments arguments{
            cutlass::gemm::GemmUniversalMode::kGemm,
            {m, n, k, 1},
            {
                a_ptr, stride_A,
                b_ptr, stride_B,
                sfa_ptr, layout_SFA,
                sfb_ptr, layout_SFB
            },
            {
                {alpha, beta},
                c_ptr, stride_C,
                d_ptr, stride_D
            }
        };
        
        // Set swizzle size for better cluster scheduling
        // arguments.scheduler.max_swizzle_size = 1;
        
        // Initialize and run GEMM
        Gemm gemm;
        size_t workspace_size = Gemm::get_workspace_size(arguments);
        
        void* workspace = nullptr;
        if (workspace_size > 0) {
            cudaMalloc(&workspace, workspace_size);
        }

        // static void* workspace = nullptr;
        // static size_t workspace_cap = 0;

        // if (workspace_size > 0) {
        //     if (workspace_size > workspace_cap) {
        //         if (workspace) cudaFree(workspace);
        //             cudaMalloc(&workspace, workspace_size);
        //             workspace_cap = workspace_size;
        //     }
        // } else {
        //     workspace = nullptr;
        // }

        cutlass::Status status = gemm.can_implement(arguments);
        if (status != cutlass::Status::kSuccess) {
            throw std::runtime_error("GEMM kernel cannot implement the given arguments");
        }
        
        status = gemm.initialize(arguments, workspace);
        if (status != cutlass::Status::kSuccess) {
            if (workspace) cudaFree(workspace);
            throw std::runtime_error("GEMM initialization failed");
        }
        
        status = gemm.run();
        if (status != cutlass::Status::kSuccess) {
            if (workspace) cudaFree(workspace);
            throw std::runtime_error("GEMM execution failed");
        }
        
        if (workspace) {
            cudaFree(workspace);
        }
    }
    
    // cudaDeviceSynchronize();
    
    return C;
}
"""

# Build the module
nvfp4_gemm_module = load_inline(
    name="nvfp4_gemm_cutlass",
    cpp_sources=[gemm_cpp],
    cuda_sources=[gemm_cuda],
    functions=["cuda_nvfp4_gemm"],
    extra_cflags=["-std=c++17"],
    extra_cuda_cflags=[
        "-std=c++17",
        # f"-I{CUTLASS_DIR}/include",
        # f"-I{CUTLASS_DIR}/tools/util/include",
        "-gencode=arch=compute_100a,code=sm_100a",
        "--ptxas-options=--gpu-name=sm_100a",
        "-O3",
        "-w",
        "--expt-relaxed-constexpr",
        "--use_fast_math",
        "-allow-unsupported-compiler",
    ],
    extra_ldflags=["-lcuda"],
    verbose=True,
)


def custom_kernel(data: input_t) -> output_t:
    """
    Wrapper function that calls the CUTLASS GEMM kernel.
    
    Input tuple structure (7 elements):
        data[0]: a - Input matrix A [m, k/2, l] in float4_e2m1fn_x2
        data[1]: b - Input matrix B [n, k/2, l] in float4_e2m1fn_x2
        data[2]: sfa_ref_cpu - Scale factors for A [m, k//16, l] in float8_e4m3fn (simple format)
        data[3]: sfb_ref_cpu - Scale factors for B [n, k//16, l] in float8_e4m3fn (simple format)
        data[4]: sfa_ref_permuted - Scale factors for A [32, 4, rest_m, 4, rest_k, l] (CUTLASS layout)
        data[5]: sfb_ref_permuted - Scale factors for B [32, 4, rest_n, 4, rest_k, l] (CUTLASS layout)
        data[6]: c - Output matrix [m, n, l] in float16 (preallocated)
    
    Returns:
        c - The output matrix in float16
    """
    # Use the permuted scale factors (data[4] and data[5]) for CUTLASS
    return nvfp4_gemm_module.cuda_nvfp4_gemm(
        data[0],  # A
        data[1],  # B
        data[4],  # SFA (permuted/CUTLASS layout)
        data[5],  # SFB (permuted/CUTLASS layout)
        data[6]   # C (output)
    )

scrolls · 318 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