Skip to content
KernelIndex
Search⌘K

submission 165770

shikhar · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

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

nvfp4_gemm_v23_cutlass3.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-165770?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
31.0µs
#213 of 369
2025-12-16

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:e88e834ac49fc29b0cbb5f648b678a26442ffdbe530063af0453f0c706f99ab4
license declaredunknown
license concludedunknown
authorsshikhar
imported2026-08-26

Techniques

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

clusterusing ClusterShape = Shape<_1, _1, _1>; // Single CTA (best for this problem size)
fp4NVFP4 Block-Scaled GEMM v23 - CUTLASS 3.x CollectiveBuilder
fused-epilogueusing FusionOperation = cutlass::epilogue::fusion::LinearCombination<

Kernel source

nvfp4_gemm_v23_cutlass3.py424 lines
"""
NVFP4 Block-Scaled GEMM v23 - CUTLASS 3.x CollectiveBuilder

Uses CUTLASS 3.x's CollectiveBuilder and GemmUniversalAdapter for:
1. Automatic TMA loads with pipelining
2. WGMMA tensor core operations
3. Proper scale factor layouts via Sm1xxBlkScaledConfig
4. Warp-specialized kernel design

This is the correct way to use tensor cores for FP4 GEMM on SM100!

Problem: C[M,N,L] = A[M,K,L] @ B[N,K,L].T

✅ k: 256; l: 1; m: 128; n: 256; seed: 1111                                                   │
✅ k: 7168; l: 1; m: 128; n: 1536; seed: 1111                                                 │
✅ k: 1536; l: 1; m: 128; n: 3072; seed: 1111                                                 │
✅ k: 256; l: 1; m: 256; n: 7168; seed: 1111                                                  │
✅ k: 2048; l: 1; m: 256; n: 7168; seed: 1111                                                 │
✅ k: 7168; l: 1; m: 2304; n: 4608; seed: 1111                                                │
✅ k: 2304; l: 1; m: 384; n: 7168; seed: 1111                                                 │
✅ k: 7168; l: 1; m: 512; n: 512; seed: 1111                                                  │
✅ k: 512; l: 1; m: 512; n: 4096; seed: 1111                                                  │
✅ k: 7168; l: 1; m: 512; n: 1536; seed: 1111```                                              │
                                                                                              │
## Benchmarks:                                                                                │
```                                                                                           │
k: 16384; l: 1; m: 128; n: 7168; seed: 1111                                                   │
 ⏱ 81.6 ± 0.01 µs                                                                             │
 ⚡ 81.6 µs 🐌 81.6 µs                                                                        │
                                                                                              │
k: 7168; l: 1; m: 128; n: 4096; seed: 1111                                                    │
 ⏱ 44.0 ± 0.04 µs                                                                             │
 ⚡ 43.6 µs 🐌 44.2 µs                                                                        │
                                                                                              │
k: 2048; l: 1; m: 128; n: 7168; seed: 1111                                                    │
 ⏱ 33.9 ± 0.01 µs                                                                             │
 ⚡ 33.9 µs 🐌 33.9 µs                                                                        │
```                                                                                           │
                                                                                              │
## Ranked Benchmark:                                                                          │
```                                                                                           │
k: 16384; l: 1; m: 128; n: 7168; seed: 1111                                                   │
 ⏱ 81.7 ± 0.00 µs                                                                             │
 ⚡ 81.7 µs 🐌 81.7 µs                                                                        │
                                                                                              │
k: 7168; l: 1; m: 128; n: 4096; seed: 1111                                                    │
 ⏱ 43.9 ± 0.04 µs                                                                             │
 ⚡ 43.7 µs 🐌 44.0 µs                                                                        │
                                                                                              │
k: 2048; l: 1; m: 128; n: 7168; seed: 1111                                                    │
 ⏱ 34.0 ± 0.01 µs                                                                             │
 ⚡ 34.0 µs 🐌 34.0 µs                                                                        │
"""

import os
import torch
from torch.utils.cpp_extension import load_inline
from pathlib import Path
from task import input_t, output_t

input_t = tuple
output_t = torch.Tensor

# Find CUTLASS path
# _current_dir = Path(__file__).resolve().parent
# if (_current_dir.parent / "cutlass" / "include").exists():
#     CUTLASS_PATH = str(_current_dir.parent / "cutlass" / "include")
#     CUTLASS_TOOLS_PATH = str(_current_dir.parent / "cutlass" / "tools" / "util" / "include")
# elif Path("/root/cutlass/include").exists():
#     CUTLASS_PATH = "/root/cutlass/include"
#     CUTLASS_TOOLS_PATH = "/root/cutlass/tools/util/include"
# else:
#     CUTLASS_PATH = None
#     CUTLASS_TOOLS_PATH = None

CUDA_SRC = r"""
#include <cuda.h>
#include <cuda_runtime.h>
#include <cuda_fp16.h>
#include <cuda_fp8.h>

#include <torch/extension.h>
#include <ATen/cuda/CUDAContext.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/numeric_types.h"
#include "cutlass/util/packed_stride.hpp"

using namespace cute;

/////////////////////////////////////////////////////////////////////////////////////////////////
/// GPU-side scale factor reformat kernel
/// Converts from sfa layout [M, K/16, L] to CUTLASS blocked format [L, flat]
/// This replaces the Python to_blocked_3d() function entirely on GPU
/////////////////////////////////////////////////////////////////////////////////////////////////

__global__ void reformat_scales_to_blocked(
    const uint8_t* __restrict__ in,   // sfa: [M, K/16, L] contiguous
    uint8_t* __restrict__ out,        // CUTLASS format: [L, n_row_blocks * n_col_blocks * 512]
    int M, int K_scales, int L,       // K_scales = K/16
    int n_row_blocks, int n_col_blocks
) {
    // Total elements = M * K_scales * L
    int tid = blockIdx.x * blockDim.x + threadIdx.x;
    int total = M * K_scales * L;
    if (tid >= total) return;

    // Decode input index from contiguous [M, K_scales, L] layout
    // Element at (i, j, b) where i=row, j=col, b=batch
    int b = tid % L;
    int rem = tid / L;
    int j = rem % K_scales;
    int i = rem / K_scales;

    // Compute to_blocked_3d transformation:
    // row_blk = i // 128, col_blk = j // 4
    // inner_row = i % 128, inner_col = j % 4
    // out = (row_blk * n_col_blocks + col_blk) * 512 + (inner_row % 32) * 16 + (inner_row // 32) * 4 + inner_col
    int row_blk = i / 128;
    int col_blk = j / 4;
    int inner_row = i % 128;
    int inner_col = j % 4;

    int out_idx = b * (n_row_blocks * n_col_blocks * 512) +
                  (row_blk * n_col_blocks + col_blk) * 512 +
                  (inner_row % 32) * 16 +
                  (inner_row / 32) * 4 +
                  inner_col;

    out[out_idx] = in[tid];
}

// #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)

/////////////////////////////////////////////////////////////////////////////////////////////////
/// GEMM kernel configurations
/////////////////////////////////////////////////////////////////////////////////////////////////

// A matrix configuration
using         ElementA    = cutlass::nv_float4_t<cutlass::float_e2m1_t>;    // FP4 e2m1
using         LayoutATag  = cutlass::layout::RowMajor;                      // K-major
constexpr int AlignmentA  = 32;

// B matrix configuration
using         ElementB    = cutlass::nv_float4_t<cutlass::float_e2m1_t>;    // FP4 e2m1
using         LayoutBTag  = cutlass::layout::ColumnMajor;                   // K-major (transposed)
constexpr int AlignmentB  = 32;

// C/D matrix configuration - output FP16
using         ElementC    = cutlass::half_t;
using         ElementD    = 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;

// Kernel configuration
using ElementAccumulator  = float;
using ElementCompute      = float;
using ArchTag             = cutlass::arch::Sm100;
using OperatorClass       = cutlass::arch::OpClassBlockScaledTensorOp;

// MMA tile shape - 128x128x256 for FP4 (standard config)
using MmaTileShape        = Shape<_128, _128, _256>;
using ClusterShape        = Shape<_1, _1, _1>;  // Single CTA (best for this problem size)

constexpr int ScaleFactorVectorSize = 16;  // 16 FP4 elements per scale

// Simple linear combination epilogue (no output quantization)
using FusionOperation = cutlass::epilogue::fusion::LinearCombination<
    ElementD,
    ElementCompute>;

using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
    ArchTag, OperatorClass,
    MmaTileShape, ClusterShape,
    cutlass::epilogue::collective::EpilogueTileAuto,
    ElementAccumulator, ElementCompute,
    ElementC, LayoutCTag, AlignmentC,
    ElementD, LayoutDTag, AlignmentD,
    cutlass::epilogue::collective::EpilogueScheduleAuto,
    FusionOperation
>::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))>,
    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>;

// Layout types
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;

// #endif // CUTLASS_ARCH_MMA_SM100_SUPPORTED

torch::Tensor cuda_nvfp4_gemm_v23_cutlass3(
    torch::Tensor A,         // [M, K/2, L] packed FP4
    torch::Tensor B,         // [N, K/2, L] packed FP4
    torch::Tensor SFA,       // [M, K/16, L] FP8 scales (CONTIGUOUS - no Python overhead!)
    torch::Tensor SFB,       // [N, K/16, L] FP8 scales (CONTIGUOUS - no Python overhead!)
    torch::Tensor C          // [M, N, L] FP16 output
) {
// #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
    const int M = A.size(0);
    const int K_packed = A.size(1);
    const int L = A.size(2);
    const int N = B.size(0);
    const int K = K_packed * 2;  // Full K dimension
    const int K_scales = K / 16;

    // Compute blocked layout dimensions
    const int n_row_blocks_a = (M + 127) / 128;
    const int n_col_blocks_a = (K_scales + 3) / 4;
    const int n_row_blocks_b = (N + 127) / 128;
    const int n_col_blocks_b = (K_scales + 3) / 4;

    // =========================================================================
    // GPU-side scale factor reformatting (ZERO Python overhead!)
    // Converts from [M, K/16, L] to CUTLASS blocked format [L, n_row_blocks * n_col_blocks * 512]
    // =========================================================================
    int scale_size_a = n_row_blocks_a * n_col_blocks_a * 512 * L;
    int scale_size_b = n_row_blocks_b * n_col_blocks_b * 512 * L;

    // Allocate reformatted buffers
    auto sfa_reformat = torch::empty({L, scale_size_a / L},
        torch::TensorOptions().dtype(torch::kUInt8).device(A.device()));
    auto sfb_reformat = torch::empty({L, scale_size_b / L},
        torch::TensorOptions().dtype(torch::kUInt8).device(A.device()));

    // Launch reformat kernels - one thread per input element
    constexpr int threads = 256;
    int total_a = M * K_scales * L;
    int total_b = N * K_scales * L;
    int blocks_a = (total_a + threads - 1) / threads;
    int blocks_b = (total_b + threads - 1) / threads;

    reformat_scales_to_blocked<<<blocks_a, threads>>>(
        reinterpret_cast<const uint8_t*>(SFA.data_ptr()),
        reinterpret_cast<uint8_t*>(sfa_reformat.data_ptr()),
        M, K_scales, L, n_row_blocks_a, n_col_blocks_a);

    reformat_scales_to_blocked<<<blocks_b, threads>>>(
        reinterpret_cast<const uint8_t*>(SFB.data_ptr()),
        reinterpret_cast<uint8_t*>(sfb_reformat.data_ptr()),
        N, K_scales, L, n_row_blocks_b, n_col_blocks_b);

    // =========================================================================
    // CUTLASS GEMM
    // =========================================================================

    // Create strides for matrices
    // A is [M, K, L] row-major (K contiguous)
    auto stride_A = cutlass::make_cute_packed_stride(StrideA{}, {M, K, L});
    // B is [N, K, L] but transposed in GEMM, so column-major
    auto stride_B = cutlass::make_cute_packed_stride(StrideB{}, {N, K, L});
    // C/D are [M, N, L] row-major
    auto stride_C = cutlass::make_cute_packed_stride(StrideC{}, {M, N, L});
    auto stride_D = cutlass::make_cute_packed_stride(StrideD{}, {M, N, L});

    // Create scale factor layouts using CUTLASS's built-in config
    auto layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(make_shape(M, N, K, L));
    auto layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(make_shape(M, N, K, L));

    // Create the GEMM adapter
    Gemm gemm;

    // Create arguments - use reformatted scales!
    typename Gemm::Arguments arguments{
        cutlass::gemm::GemmUniversalMode::kGemm,
        {M, N, K, L},
        {
            reinterpret_cast<ElementA::DataType const*>(A.data_ptr()),
            stride_A,
            reinterpret_cast<ElementB::DataType const*>(B.data_ptr()),
            stride_B,
            reinterpret_cast<ElementA::ScaleFactorType const*>(sfa_reformat.data_ptr()),
            layout_SFA,
            reinterpret_cast<ElementB::ScaleFactorType const*>(sfb_reformat.data_ptr()),
            layout_SFB
        },
        {
            {1.0f, 0.0f},  // alpha, beta
            reinterpret_cast<ElementC const*>(C.data_ptr()),  // C input
            stride_C,
            reinterpret_cast<ElementD*>(C.data_ptr()),  // D output (same as C for in-place)
            stride_D
        }
    };

    // Query workspace size
    size_t workspace_size = Gemm::get_workspace_size(arguments);
    auto workspace = torch::empty({static_cast<int64_t>(workspace_size)},
                                   torch::TensorOptions().dtype(torch::kUInt8).device(A.device()));

    // Check if problem size is supported
    auto status = gemm.can_implement(arguments);
    if (status != cutlass::Status::kSuccess) {
        throw std::runtime_error("CUTLASS cannot implement this problem size");
    }

    // Initialize and run
    status = gemm.initialize(arguments, workspace.data_ptr());
    if (status != cutlass::Status::kSuccess) {
        throw std::runtime_error("CUTLASS initialization failed");
    }

    status = gemm.run();
    if (status != cutlass::Status::kSuccess) {
        throw std::runtime_error("CUTLASS kernel execution failed");
    }

    return C;
// #else
//     throw std::runtime_error("CUTLASS SM100 MMA not supported");
// #endif
}
"""

CPP_SRC = r"""
#include <torch/extension.h>

torch::Tensor cuda_nvfp4_gemm_v23_cutlass3(
    torch::Tensor A,
    torch::Tensor B,
    torch::Tensor SFA,
    torch::Tensor SFB,
    torch::Tensor C
);
"""

# Build the module
# HAS_V23 = False
# nvfp4_module = None

# if CUTLASS_PATH is not None:
# try:
extra_cuda_cflags = [
    "-std=c++17",
    "-gencode=arch=compute_100a,code=sm_100a",
    "--ptxas-options=--gpu-name=sm_100a",
    "-O3",
    "-w",
    "--use_fast_math",
    "-allow-unsupported-compiler",
    # "-I/root/cutlass/include",  # CUTLASS headers on Modal
    # "-I/root/cutlass/tools/util/include",  # CUTLASS tools headers
]

nvfp4_module = load_inline(
    name="nvfp4_gemm_v23_cutlass3",
    cpp_sources=[CPP_SRC],
    cuda_sources=[CUDA_SRC],
    functions=["cuda_nvfp4_gemm_v23_cutlass3"],
    extra_cuda_cflags=extra_cuda_cflags,
    extra_ldflags=["-lcuda"],
    verbose=False,
)
# HAS_V23 = True  # For test harness detection
# except Exception as e:
    # HAS_V23 = False
    # nvfp4_module = None
    # print(f"v23 compilation failed: {e}")
# else:
#     # print("CUTLASS path not found, v23 not available")


def ceil_div(a, b):
    """Ceiling division helper."""
    return (a + b - 1) // b


def custom_kernel(data: input_t) -> output_t:
    """
    NVFP4 block-scaled GEMM v23 - CUTLASS 3.x CollectiveBuilder.

    Uses proper tensor core operations via GemmUniversalAdapter.
    GPU-side scale reformatting - ZERO Python overhead!

    Uses original sfa/sfb tensors directly (already contiguous).
    GPU kernel converts from [M, K/16, L] to CUTLASS blocked format.
    """
    a, b, sfa, sfb, sfa_perm, sfb_perm, c = data

    if not a.is_cuda:
        a = a.cuda()
    if not b.is_cuda:
        b = b.cuda()
    if not sfa.is_cuda:
        sfa = sfa.cuda()
    if not sfb.is_cuda:
        sfb = sfb.cuda()
    if not c.is_cuda:
        c = c.cuda()

    # NO .contiguous() needed - sfa/sfb are already contiguous from ref.py!
    # GPU kernel handles the to_blocked_3d transformation entirely on GPU.
    return nvfp4_module.cuda_nvfp4_gemm_v23_cutlass3(a, b, sfa, sfb, c)
scrolls · 424 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