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
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.
cluster
using ClusterShape = Shape<_1, _1, _1>; // Single CTA (best for this problem size)fp4
NVFP4 Block-Scaled GEMM v23 - CUTLASS 3.x CollectiveBuilderfused-epilogue
using 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