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
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.
cluster
using ClusterShape = Shape<_1,_2,_1>;fp4
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;fused-epilogue
using 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