Skip to content
KernelIndex
Search⌘K

submission 133128

IncompetentGeometer · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

submission_cutlass.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-133128?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
38.0µs
#235 of 369
2025-12-08

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:dbb724b3e732322624b6f3ef33539b7a2069de0a01d9c9ca88ea02b3e972a46c
license declaredunknown
license concludedunknown
authorsIncompetentGeometer
imported2026-08-26

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

clusterusing ClusterShape = Shape<_1, _1, _1>;
fp4CUTLASS NVFP4 Block-Scaled GEMM Submission
fp8- SFA/SFB: Block scale factors in FP8 (e4m3)
fused-epilogueusing CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<

Kernel source

submission_cutlass.py267 lines
"""
CUTLASS NVFP4 Block-Scaled GEMM Submission

Adapts the blackwell_narrow_precision_gemm.cu CUTLASS example
to the competition submission format.

Computes: D = (A * SFA) @ (B * SFB)^T
Where:
  - A: [M, K] in NVFP4 (float_e2m1)
  - B: [N, K] in NVFP4 (float_e2m1)  
  - SFA/SFB: Block scale factors in FP8 (e4m3)
  - D: [M, N] in FP16
"""
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t

# C++ function declarations
CPP_SRC = r"""
#include <torch/extension.h>

torch::Tensor cutlass_nvfp4_gemm(
    torch::Tensor A,
    torch::Tensor B, 
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor D,
    int M, int N, int K
);
"""

# CUDA source with CUTLASS kernel
CUDA_SRC = r"""
#include <torch/extension.h>
#include <cuda.h>
#include <cuda_runtime.h>

#include "cutlass/cutlass.h"
#include "cute/tensor.hpp"
#include "cutlass/detail/sm100_blockscaled_layout.hpp"
#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/util/packed_stride.hpp"

using namespace cute;

// ============================================================================
// GEMM Kernel Configuration
// ============================================================================

// A matrix: NVFP4 (e2m1) with block scaling
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutATag = cutlass::layout::RowMajor;  // A is [M, K]
constexpr int AlignmentA = 32;

// B matrix: NVFP4 (e2m1) with block scaling  
using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutBTag = cutlass::layout::ColumnMajor;  // B is [N, K], transposed for TN layout
constexpr int AlignmentB = 32;

// Output: FP16 (competition requirement)
using ElementD = cutlass::half_t;
using ElementC = cutlass::half_t;
using LayoutCTag = cutlass::layout::RowMajor;
using LayoutDTag = cutlass::layout::RowMajor;
constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;

// Accumulator and compute types
using ElementAccumulator = float;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;

// Tile configuration - using smaller tiles for initial testing
using MmaTileShape = Shape<_128, _128, _256>;
using ClusterShape = Shape<_1, _1, _1>;

// ============================================================================
// Build Epilogue 
// ============================================================================

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;

// ============================================================================
// Build Mainloop
// ============================================================================

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))
    >,
    cutlass::gemm::collective::KernelScheduleAuto
>::CollectiveOp;

// ============================================================================
// Build GEMM Kernel
// ============================================================================

using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
    Shape<int, int, int, int>,  // Problem shape: M, N, K, L
    CollectiveMainloop,
    CollectiveEpilogue,
    void  // No tile scheduler override
>;

using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;

// Type aliases for strides and layouts
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;

// ============================================================================
// Host Wrapper Function
// ============================================================================

torch::Tensor cutlass_nvfp4_gemm(
    torch::Tensor A,      // [M, K/2, L] in float4_e2m1fn_x2 (packed)
    torch::Tensor B,      // [N, K/2, L] in float4_e2m1fn_x2 (packed)
    torch::Tensor SFA,    // [32, 4, ceil(M/128), 4, ceil(K/64), L] permuted
    torch::Tensor SFB,    // [32, 4, ceil(N/128), 4, ceil(K/64), L] permuted
    torch::Tensor D,      // [M, N, L] output
    int M, int N, int K
) {
    // Get batch count from output tensor
    int L = D.size(2);
    
    // Create strides
    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});
    
    // Create scale factor layouts using CUTLASS's helper
    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)
    );
    
    // Get data pointers - reinterpret_cast to CUTLASS types
    auto* ptr_A = reinterpret_cast<ElementA::DataType*>(A.data_ptr());
    auto* ptr_B = reinterpret_cast<ElementB::DataType*>(B.data_ptr());
    auto* ptr_SFA = reinterpret_cast<ElementA::ScaleFactorType*>(SFA.data_ptr());
    auto* ptr_SFB = reinterpret_cast<ElementB::ScaleFactorType*>(SFB.data_ptr());
    auto* ptr_C = reinterpret_cast<ElementC*>(D.data_ptr());  // C = D initially (beta=0)
    auto* ptr_D = reinterpret_cast<ElementD*>(D.data_ptr());
    
    // Create GEMM arguments
    typename Gemm::Arguments arguments{
        cutlass::gemm::GemmUniversalMode::kGemm,
        {M, N, K, L},  // Problem shape
        {   // Mainloop arguments
            ptr_A, stride_A,
            ptr_B, stride_B,
            ptr_SFA, layout_SFA,
            ptr_SFB, layout_SFB
        },
        {   // Epilogue arguments
            {1.0f, 0.0f},  // alpha=1, beta=0
            ptr_C, stride_C,
            ptr_D, stride_D
        }
    };
    
    // Instantiate GEMM
    Gemm gemm;
    
    // Query workspace size
    size_t workspace_size = Gemm::get_workspace_size(arguments);
    
    // Allocate workspace
    auto workspace = torch::empty({static_cast<long>(workspace_size)}, 
                                   torch::TensorOptions().dtype(torch::kUInt8).device(A.device()));
    
    // Check if this problem size is supported
    auto status = gemm.can_implement(arguments);
    if (status != cutlass::Status::kSuccess) {
        throw std::runtime_error("CUTLASS cannot implement this GEMM configuration");
    }
    
    // Initialize
    status = gemm.initialize(arguments, workspace.data_ptr<uint8_t>());
    if (status != cutlass::Status::kSuccess) {
        throw std::runtime_error("CUTLASS initialization failed");
    }
    
    // Run the GEMM
    status = gemm.run();
    if (status != cutlass::Status::kSuccess) {
        throw std::runtime_error("CUTLASS GEMM execution failed");
    }
    
    // Synchronize
    cudaDeviceSynchronize();
    
    return D;
}
"""

# Compile CUDA extension
CUDA_FLAGS = [
    "-std=c++17",
    "-gencode=arch=compute_100a,code=sm_100a",
    "--ptxas-options=--gpu-name=sm_100a",
    "-O3",
    "-w",  # Suppress warnings
    "--use_fast_math",
]

LD_FLAGS = [
    "-lcuda",
    "-lcublas",
]

# Compile the module
nvfp4_gemm_module = load_inline(
    name="cutlass_nvfp4_gemm",
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=["cutlass_nvfp4_gemm"],
    extra_cuda_cflags=CUDA_FLAGS,
    extra_ldflags=LD_FLAGS,
    verbose=True,
)


def custom_kernel(data: input_t) -> output_t:
    """
    Competition entry point.
    
    data = (a_ref, b_ref, sfa_ref_cpu, sfb_ref_cpu, sfa_permuted, sfb_permuted, c_ref)
    
    We use the permuted scale factors (indices 4, 5) as they match CUTLASS's expected layout.
    """
    a_ref, b_ref, _, _, sfa_perm, sfb_perm, c_ref = data
    
    # Get dimensions
    M = a_ref.size(0)
    K = a_ref.size(1) * 2  # K/2 is stored (packed FP4 pairs)
    N = b_ref.size(0)
    
    # Call CUTLASS kernel
    return nvfp4_gemm_module.cutlass_nvfp4_gemm(
        a_ref, b_ref, sfa_perm, sfb_perm, c_ref, M, N, K
    )
scrolls · 267 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