submission 476965
J · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 979 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-476965?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:90d73ab5ae42555af2f6d135e488b3859156c7fa6d94228734ec31c6f69f7aec
license declaredunknown
license concludedunknown
authorsJ
imported2026-08-15
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
using ClusterShape = Shape<int32_t,int32_t,_1>;fp4
using ElementInput = cutlass::float_e2m1_t;fused-epilogue
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;warp-specialization
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;Kernel source
submission.py979 lines
#!POPCORN leaderboard nvfp4_group_gemm
#!POPCORN gpu NVIDIA
from typing import Dict, List, Tuple
import weakref
import sys
import torch
import torch.utils.cpp_extension
from task import input_t, output_t
from utils import make_match_reference
sf_vec_size = 16
def ceil_div(a, b):
return (a + b - 1) // b
def to_blocked(input_matrix: torch.Tensor) -> torch.Tensor:
rows, cols = input_matrix.shape
n_row_blocks = ceil_div(rows, 128)
n_col_blocks = ceil_div(cols, 4)
padded_rows = n_row_blocks * 128
padded_cols = n_col_blocks * 4
if padded_rows != rows or padded_cols != cols:
padded = torch.nn.functional.pad(
input_matrix,
(0, padded_cols - cols, 0, padded_rows - rows),
mode="constant",
value=0,
)
else:
padded = input_matrix
blocks = padded.view(n_row_blocks, 128, n_col_blocks, 4).permute(0, 2, 1, 3)
rearranged = blocks.reshape(-1, 4, 32, 4).transpose(1, 2).reshape(-1, 32, 16)
return rearranged.flatten()
def _unpack_input(data: input_t):
if len(data) == 3:
abc_tensors, sfasfb_tensors, problem_sizes = data
return abc_tensors, sfasfb_tensors, None, problem_sizes
if len(data) == 4:
abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = data
return abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes
raise ValueError(f"Unexpected input format with {len(data)} elements")
CUDA_SOURCE = r"""
#include "cutlass/cutlass.h"
#include "cute/tensor.hpp"
#include "cutlass/tensor_ref.h"
#include "cutlass/epilogue/collective/default_epilogue.hpp"
#include "cutlass/epilogue/thread/linear_combination.h"
#include "cutlass/gemm/dispatch_policy.hpp"
#include "cutlass/gemm/group_array_problem_shape.hpp"
#include "cutlass/gemm/collective/collective_builder.hpp"
#include "cutlass/epilogue/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;
using ProblemShape = cutlass::gemm::GroupProblemShape<Shape<int,int,int>>;
using ElementInput = cutlass::float_e2m1_t;
using ElementSF = cutlass::float_ue4m3_t;
using ElementC = cutlass::half_t;
using ElementA = cutlass::nv_float4_t<ElementInput>;
using LayoutA = cutlass::layout::RowMajor;
constexpr int AlignmentA = 32;
using ElementB = cutlass::nv_float4_t<ElementInput>;
using LayoutB = cutlass::layout::ColumnMajor;
constexpr int AlignmentB = 32;
using ElementD = ElementC;
using LayoutC = cutlass::layout::RowMajor;
constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
constexpr int AlignmentD = 128 / cutlass::sizeof_bits<ElementD>::value;
using ElementAccumulator = float;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp;
using StageCountType = cutlass::gemm::collective::StageCountAuto;
using ClusterShape = Shape<int32_t,int32_t,_1>;
struct MMA1SMConfig {
using MmaTileShape = Shape<_128,_256,_256>;
using KernelSchedule = cutlass::gemm::KernelPtrArrayTmaWarpSpecialized1SmNvf4Sm100;
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;
};
struct MMA2SMConfig {
using MmaTileShape = Shape<_256,_256,_256>;
using KernelSchedule = cutlass::gemm::KernelPtrArrayTmaWarpSpecialized2SmNvf4Sm100;
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized2Sm;
};
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag,
OperatorClass,
typename MMA1SMConfig::MmaTileShape,
ClusterShape,
Shape<_128,_64>,
ElementAccumulator,
ElementAccumulator,
ElementC,
LayoutC *,
AlignmentC,
ElementD,
LayoutC *,
AlignmentD,
typename MMA1SMConfig::EpilogueSchedule
>::CollectiveOp;
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
ArchTag,
OperatorClass,
ElementA,
LayoutA *,
AlignmentA,
ElementB,
LayoutB *,
AlignmentB,
ElementAccumulator,
typename MMA1SMConfig::MmaTileShape,
ClusterShape,
cutlass::gemm::collective::StageCountAutoCarveout<
static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
typename MMA1SMConfig::KernelSchedule
>::CollectiveOp;
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
ProblemShape,
CollectiveMainloop,
CollectiveEpilogue
>;
using Gemm1SM = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
using CollectiveEpilogue2SM = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag,
OperatorClass,
typename MMA2SMConfig::MmaTileShape,
ClusterShape,
Shape<_128,_64>,
ElementAccumulator,
ElementAccumulator,
ElementC,
LayoutC *,
AlignmentC,
ElementD,
LayoutC *,
AlignmentD,
typename MMA2SMConfig::EpilogueSchedule
>::CollectiveOp;
using CollectiveMainloop2SM = typename cutlass::gemm::collective::CollectiveBuilder<
ArchTag,
OperatorClass,
ElementA,
LayoutA *,
AlignmentA,
ElementB,
LayoutB *,
AlignmentB,
ElementAccumulator,
typename MMA2SMConfig::MmaTileShape,
ClusterShape,
cutlass::gemm::collective::StageCountAutoCarveout<
static_cast<int>(sizeof(typename CollectiveEpilogue2SM::SharedStorage))>,
typename MMA2SMConfig::KernelSchedule
>::CollectiveOp;
using GemmKernel2SM = cutlass::gemm::kernel::GemmUniversal<
ProblemShape,
CollectiveMainloop2SM,
CollectiveEpilogue2SM
>;
using Gemm2SM = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel2SM>;
using Gemm = Gemm1SM;
using StrideA = typename Gemm::GemmKernel::InternalStrideA;
using StrideB = typename Gemm::GemmKernel::InternalStrideB;
using StrideC = typename Gemm::GemmKernel::InternalStrideC;
using StrideD = typename Gemm::GemmKernel::InternalStrideD;
using ArrayElementA = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementA;
using ArrayElementB = typename Gemm::GemmKernel::CollectiveMainloop::ArrayElementB;
using InternalLayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFA;
using InternalLayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFB;
using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
static_assert(sizeof(typename Gemm1SM::GemmKernel::InternalStrideA) == sizeof(typename Gemm2SM::GemmKernel::InternalStrideA));
static_assert(sizeof(typename Gemm1SM::GemmKernel::InternalStrideB) == sizeof(typename Gemm2SM::GemmKernel::InternalStrideB));
static_assert(sizeof(typename Gemm1SM::GemmKernel::InternalStrideC) == sizeof(typename Gemm2SM::GemmKernel::InternalStrideC));
static_assert(sizeof(typename Gemm1SM::GemmKernel::InternalStrideD) == sizeof(typename Gemm2SM::GemmKernel::InternalStrideD));
static_assert(sizeof(typename Gemm1SM::GemmKernel::CollectiveMainloop::InternalLayoutSFA) == sizeof(typename Gemm2SM::GemmKernel::CollectiveMainloop::InternalLayoutSFA));
static_assert(sizeof(typename Gemm1SM::GemmKernel::CollectiveMainloop::InternalLayoutSFB) == sizeof(typename Gemm2SM::GemmKernel::CollectiveMainloop::InternalLayoutSFB));
template <typename GemmT>
int run_cutlass_grouped_gemm_impl(
int num_groups,
void* d_problem_sizes,
void* d_ptr_A,
void* d_stride_A,
void* d_ptr_B,
void* d_stride_B,
void* d_ptr_SFA,
void* d_layout_SFA,
void* d_ptr_SFB,
void* d_layout_SFB,
void* d_ptr_C,
void* d_stride_C,
void* d_ptr_D,
void* d_stride_D,
void* d_workspace,
size_t workspace_size,
void* h_problem_shapes,
int cluster_x,
int cluster_fallback_x,
bool run_kernel
) {
using StrideA_t = typename GemmT::GemmKernel::InternalStrideA;
using StrideB_t = typename GemmT::GemmKernel::InternalStrideB;
using StrideC_t = typename GemmT::GemmKernel::InternalStrideC;
using StrideD_t = typename GemmT::GemmKernel::InternalStrideD;
using ArrayElementA_t = typename GemmT::GemmKernel::CollectiveMainloop::ArrayElementA;
using ArrayElementB_t = typename GemmT::GemmKernel::CollectiveMainloop::ArrayElementB;
using InternalLayoutSFA_t = typename GemmT::GemmKernel::CollectiveMainloop::InternalLayoutSFA;
using InternalLayoutSFB_t = typename GemmT::GemmKernel::CollectiveMainloop::InternalLayoutSFB;
cutlass::KernelHardwareInfo hw_info;
hw_info.device_id = 0;
hw_info.sm_count = cutlass::KernelHardwareInfo::query_device_multiprocessor_count(hw_info.device_id);
hw_info.cluster_shape = dim3(cluster_x, 1, 1);
hw_info.cluster_shape_fallback = dim3(cluster_fallback_x, 1, 1);
typename GemmT::GemmKernel::TileSchedulerArguments scheduler;
typename GemmT::Arguments arguments;
decltype(arguments.epilogue.thread) fusion_args;
fusion_args.alpha_ptr = nullptr;
fusion_args.beta_ptr = nullptr;
fusion_args.alpha = 1.0f;
fusion_args.beta = 0.0f;
fusion_args.alpha_ptr_array = nullptr;
fusion_args.beta_ptr_array = nullptr;
fusion_args.dAlpha = {_0{}, _0{}, 0};
fusion_args.dBeta = {_0{}, _0{}, 0};
arguments = typename GemmT::Arguments {
cutlass::gemm::GemmUniversalMode::kGrouped,
{num_groups,
reinterpret_cast<typename ProblemShape::UnderlyingProblemShape*>(d_problem_sizes),
reinterpret_cast<typename ProblemShape::UnderlyingProblemShape*>(h_problem_shapes)},
{reinterpret_cast<const ArrayElementA_t**>(d_ptr_A),
reinterpret_cast<StrideA_t*>(d_stride_A),
reinterpret_cast<const ArrayElementB_t**>(d_ptr_B),
reinterpret_cast<StrideB_t*>(d_stride_B),
reinterpret_cast<const ElementSF**>(d_ptr_SFA),
reinterpret_cast<InternalLayoutSFA_t*>(d_layout_SFA),
reinterpret_cast<const ElementSF**>(d_ptr_SFB),
reinterpret_cast<InternalLayoutSFB_t*>(d_layout_SFB)},
{fusion_args,
reinterpret_cast<const ElementC**>(d_ptr_C),
reinterpret_cast<StrideC_t*>(d_stride_C),
reinterpret_cast<ElementD**>(d_ptr_D),
reinterpret_cast<StrideD_t*>(d_stride_D)},
hw_info, scheduler
};
size_t needed_ws = GemmT::get_workspace_size(arguments);
if (needed_ws > workspace_size) {
return -1;
}
static GemmT gemm_op;
auto status = gemm_op.can_implement(arguments);
if (status != cutlass::Status::kSuccess) {
return -2;
}
if (!run_kernel) {
return 0;
}
status = gemm_op.initialize(arguments, d_workspace, nullptr);
if (status != cutlass::Status::kSuccess) {
return -3;
}
status = gemm_op.run(nullptr, nullptr, true);
if (status != cutlass::Status::kSuccess) {
return -4;
}
return 0;
}
extern "C" int run_cutlass_grouped_gemm_mode(
int num_groups,
void* d_problem_sizes,
void* d_ptr_A,
void* d_stride_A,
void* d_ptr_B,
void* d_stride_B,
void* d_ptr_SFA,
void* d_layout_SFA,
void* d_ptr_SFB,
void* d_layout_SFB,
void* d_ptr_C,
void* d_stride_C,
void* d_ptr_D,
void* d_stride_D,
void* d_workspace,
size_t workspace_size,
void* h_problem_shapes,
int mode,
bool run_kernel
) {
if (mode == 1) {
return run_cutlass_grouped_gemm_impl<Gemm1SM>(
num_groups,
d_problem_sizes,
d_ptr_A,
d_stride_A,
d_ptr_B,
d_stride_B,
d_ptr_SFA,
d_layout_SFA,
d_ptr_SFB,
d_layout_SFB,
d_ptr_C,
d_stride_C,
d_ptr_D,
d_stride_D,
d_workspace,
workspace_size,
h_problem_shapes,
1,
1,
run_kernel
);
}
if (mode == 2) {
return run_cutlass_grouped_gemm_impl<Gemm2SM>(
num_groups,
d_problem_sizes,
d_ptr_A,
d_stride_A,
d_ptr_B,
d_stride_B,
d_ptr_SFA,
d_layout_SFA,
d_ptr_SFB,
d_layout_SFB,
d_ptr_C,
d_stride_C,
d_ptr_D,
d_stride_D,
d_workspace,
workspace_size,
h_problem_shapes,
2,
2,
run_kernel
);
}
return -9;
}
// Helper to get sizes of internal types
extern "C" void get_type_sizes(int* sizes) {
sizes[0] = sizeof(StrideA);
sizes[1] = sizeof(StrideB);
sizes[2] = sizeof(StrideC);
sizes[3] = sizeof(StrideD);
sizes[4] = sizeof(InternalLayoutSFA);
sizes[5] = sizeof(InternalLayoutSFB);
sizes[6] = sizeof(cute::Shape<int,int,int>);
sizes[7] = sizeof(ElementA);
sizes[8] = sizeof(ArrayElementA);
sizes[9] = sizeof(ElementB);
sizes[10] = sizeof(ArrayElementB);
sizes[11] = sizeof(ElementSF);
}
template <typename GemmT>
void fill_metadata_impl(
int num_groups,
int* mnk,
void* out_problem_shapes,
void* out_stride_A,
void* out_stride_B,
void* out_stride_C,
void* out_stride_D,
void* out_layout_SFA,
void* out_layout_SFB
) {
using StrideA_t = typename GemmT::GemmKernel::InternalStrideA;
using StrideB_t = typename GemmT::GemmKernel::InternalStrideB;
using StrideC_t = typename GemmT::GemmKernel::InternalStrideC;
using StrideD_t = typename GemmT::GemmKernel::InternalStrideD;
using InternalLayoutSFA_t = typename GemmT::GemmKernel::CollectiveMainloop::InternalLayoutSFA;
using InternalLayoutSFB_t = typename GemmT::GemmKernel::CollectiveMainloop::InternalLayoutSFB;
using Sm1xxBlkScaledConfig_t = typename GemmT::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
auto* ps = reinterpret_cast<cute::Shape<int,int,int>*>(out_problem_shapes);
auto* sa = reinterpret_cast<StrideA_t*>(out_stride_A);
auto* sb = reinterpret_cast<StrideB_t*>(out_stride_B);
auto* sc = reinterpret_cast<StrideC_t*>(out_stride_C);
auto* sd = reinterpret_cast<StrideD_t*>(out_stride_D);
auto* lsfa = reinterpret_cast<InternalLayoutSFA_t*>(out_layout_SFA);
auto* lsfb = reinterpret_cast<InternalLayoutSFB_t*>(out_layout_SFB);
for (int i = 0; i < num_groups; i++) {
int M = mnk[i*3+0], N = mnk[i*3+1], K = mnk[i*3+2];
ps[i] = cute::make_shape(M, N, K);
sa[i] = cutlass::make_cute_packed_stride(StrideA_t{}, {M, K, 1});
sb[i] = cutlass::make_cute_packed_stride(StrideB_t{}, {N, K, 1});
sc[i] = cutlass::make_cute_packed_stride(StrideC_t{}, {M, N, 1});
sd[i] = cutlass::make_cute_packed_stride(StrideD_t{}, {M, N, 1});
lsfa[i] = Sm1xxBlkScaledConfig_t::tile_atom_to_shape_SFA(cute::make_shape(M, N, K, 1));
lsfb[i] = Sm1xxBlkScaledConfig_t::tile_atom_to_shape_SFB(cute::make_shape(M, N, K, 1));
}
}
extern "C" int fill_metadata_mode(
int mode,
int num_groups,
int* mnk,
void* out_problem_shapes,
void* out_stride_A,
void* out_stride_B,
void* out_stride_C,
void* out_stride_D,
void* out_layout_SFA,
void* out_layout_SFB
) {
if (mode == 1) {
fill_metadata_impl<Gemm1SM>(
num_groups,
mnk,
out_problem_shapes,
out_stride_A,
out_stride_B,
out_stride_C,
out_stride_D,
out_layout_SFA,
out_layout_SFB
);
return 0;
}
if (mode == 2) {
fill_metadata_impl<Gemm2SM>(
num_groups,
mnk,
out_problem_shapes,
out_stride_A,
out_stride_B,
out_stride_C,
out_stride_D,
out_layout_SFA,
out_layout_SFB
);
return 0;
}
return -1;
}
"""
CPP_SOURCE = r"""
#include <torch/extension.h>
#include <ATen/ATen.h>
#include <vector>
extern "C" int run_cutlass_grouped_gemm_mode(
int num_groups,
void* d_problem_sizes,
void* d_ptr_A, void* d_stride_A,
void* d_ptr_B, void* d_stride_B,
void* d_ptr_SFA, void* d_layout_SFA,
void* d_ptr_SFB, void* d_layout_SFB,
void* d_ptr_C, void* d_stride_C,
void* d_ptr_D, void* d_stride_D,
void* d_workspace, size_t workspace_size,
void* h_problem_shapes,
int mode,
bool run_kernel
);
extern "C" void get_type_sizes(int* sizes);
extern "C" int fill_metadata_mode(
int mode,
int num_groups,
int* mnk,
void* out_problem_shapes,
void* out_stride_A,
void* out_stride_B,
void* out_stride_C,
void* out_stride_D,
void* out_layout_SFA,
void* out_layout_SFB
);
std::vector<int> get_sizes() {
std::vector<int> sizes(12);
get_type_sizes(sizes.data());
return sizes;
}
int run_grouped(
int num_groups,
at::Tensor d_problem_sizes,
at::Tensor d_ptr_A, at::Tensor d_stride_A,
at::Tensor d_ptr_B, at::Tensor d_stride_B,
at::Tensor d_ptr_SFA, at::Tensor d_layout_SFA,
at::Tensor d_ptr_SFB, at::Tensor d_layout_SFB,
at::Tensor d_ptr_C, at::Tensor d_stride_C,
at::Tensor d_ptr_D, at::Tensor d_stride_D,
at::Tensor d_workspace,
at::Tensor h_problem_shapes,
int mode
) {
return run_cutlass_grouped_gemm_mode(
num_groups,
d_problem_sizes.data_ptr(),
d_ptr_A.data_ptr(), d_stride_A.data_ptr(),
d_ptr_B.data_ptr(), d_stride_B.data_ptr(),
d_ptr_SFA.data_ptr(), d_layout_SFA.data_ptr(),
d_ptr_SFB.data_ptr(), d_layout_SFB.data_ptr(),
d_ptr_C.data_ptr(), d_stride_C.data_ptr(),
d_ptr_D.data_ptr(), d_stride_D.data_ptr(),
d_workspace.data_ptr(), d_workspace.nbytes(),
h_problem_shapes.data_ptr(),
mode,
true
);
}
int can_implement_grouped(
int num_groups,
at::Tensor d_problem_sizes,
at::Tensor d_ptr_A, at::Tensor d_stride_A,
at::Tensor d_ptr_B, at::Tensor d_stride_B,
at::Tensor d_ptr_SFA, at::Tensor d_layout_SFA,
at::Tensor d_ptr_SFB, at::Tensor d_layout_SFB,
at::Tensor d_ptr_C, at::Tensor d_stride_C,
at::Tensor d_ptr_D, at::Tensor d_stride_D,
at::Tensor d_workspace,
at::Tensor h_problem_shapes,
int mode
) {
return run_cutlass_grouped_gemm_mode(
num_groups,
d_problem_sizes.data_ptr(),
d_ptr_A.data_ptr(), d_stride_A.data_ptr(),
d_ptr_B.data_ptr(), d_stride_B.data_ptr(),
d_ptr_SFA.data_ptr(), d_layout_SFA.data_ptr(),
d_ptr_SFB.data_ptr(), d_layout_SFB.data_ptr(),
d_ptr_C.data_ptr(), d_stride_C.data_ptr(),
d_ptr_D.data_ptr(), d_stride_D.data_ptr(),
d_workspace.data_ptr(), d_workspace.nbytes(),
h_problem_shapes.data_ptr(),
mode,
false
);
}
int fill_meta(
int mode,
int num_groups,
at::Tensor mnk,
at::Tensor ps,
at::Tensor sa,
at::Tensor sb,
at::Tensor sc,
at::Tensor sd,
at::Tensor lsfa,
at::Tensor lsfb
) {
return fill_metadata_mode(
mode,
num_groups,
mnk.data_ptr<int>(),
ps.data_ptr(),
sa.data_ptr(),
sb.data_ptr(),
sc.data_ptr(),
sd.data_ptr(),
lsfa.data_ptr(),
lsfb.data_ptr()
);
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("get_sizes", &get_sizes, "Get internal type sizes");
m.def("run", &run_grouped, "CUTLASS grouped FP4 GEMM");
m.def("can_implement", &can_implement_grouped, "CUTLASS grouped FP4 GEMM can_implement");
m.def("fill_meta", &fill_meta, "Fill metadata arrays");
}
"""
print("[BUILD] Starting CUTLASS grouped GEMM build...", file=sys.stderr, flush=True)
_MODULE = torch.utils.cpp_extension.load_inline(
name="nvfp4_group_gemm_cutlass_v2",
cpp_sources=CPP_SOURCE,
cuda_sources=CUDA_SOURCE,
extra_cuda_cflags=[
"-O3",
"--use_fast_math",
"-gencode=arch=compute_100a,code=sm_100a",
"-I/opt/cutlass/4.3.0/include",
"-I/opt/cutlass/4.3.0/build/include",
"-std=c++17",
"--expt-relaxed-constexpr",
],
extra_cflags=["-O3"],
extra_include_paths=["/usr/local/cuda/include"],
verbose=True,
)
print("[BUILD] CUTLASS build complete!", file=sys.stderr, flush=True)
_TYPE_SIZES = _MODULE.get_sizes()
print(f"[INFO] Type sizes: StrideA={_TYPE_SIZES[0]}, StrideB={_TYPE_SIZES[1]}, StrideC={_TYPE_SIZES[2]}, StrideD={_TYPE_SIZES[3]}, LayoutSFA={_TYPE_SIZES[4]}, LayoutSFB={_TYPE_SIZES[5]}, ProblemShape={_TYPE_SIZES[6]}, ElementA={_TYPE_SIZES[7]}, ArrayElementA={_TYPE_SIZES[8]}, ElementB={_TYPE_SIZES[9]}, ArrayElementB={_TYPE_SIZES[10]}, ElementSF={_TYPE_SIZES[11]}", file=sys.stderr, flush=True)
_WORKSPACE = torch.empty(32 * 1024 * 1024, dtype=torch.uint8, device="cuda")
class ScaleCache:
def __init__(self):
self._cache: Dict[Tuple[int, int, int, str], Tuple[weakref.ReferenceType[torch.Tensor], torch.Tensor]] = {}
@staticmethod
def _key(scale_3d: torch.Tensor, l_idx: int, target_device: torch.device) -> Tuple[int, int, int, str]:
return (id(scale_3d), l_idx, scale_3d._version, str(target_device))
def get_or_create(self, scale_3d: torch.Tensor, l_idx: int, target_device: torch.device) -> torch.Tensor:
key = self._key(scale_3d, l_idx, target_device)
cached = self._cache.get(key)
if cached is not None:
tensor_ref, blocked = cached
if tensor_ref() is scale_3d:
return blocked
blocked = to_blocked(scale_3d[:, :, l_idx]).to(target_device).contiguous()
self._cache[key] = (weakref.ref(scale_3d), blocked)
return blocked
_SCALE_CACHE = ScaleCache()
class ReorderedScaleCache:
def __init__(self):
self._cache: Dict[Tuple[int, int, int, str], Tuple[weakref.ReferenceType[torch.Tensor], torch.Tensor]] = {}
@staticmethod
def _key(scale_6d: torch.Tensor, l_idx: int, target_device: torch.device) -> Tuple[int, int, int, str]:
return (id(scale_6d), l_idx, scale_6d._version, str(target_device))
def get_or_create(self, scale_6d: torch.Tensor, l_idx: int, target_device: torch.device) -> torch.Tensor:
key = self._key(scale_6d, l_idx, target_device)
cached = self._cache.get(key)
if cached is not None:
tensor_ref, blocked = cached
if tensor_ref() is scale_6d:
return blocked
s = scale_6d[..., l_idx]
if s.device != target_device:
s = s.to(target_device)
blocked = s.permute(2, 4, 0, 1, 3).contiguous().view(-1)
self._cache[key] = (weakref.ref(scale_6d), blocked)
return blocked
_REORDERED_SCALE_CACHE = ReorderedScaleCache()
def _build_cutlass_metadata(problem_sizes_list, a_list, b_list, sfa_list, sfb_list, out_list, mode: int):
ng = len(problem_sizes_list)
sz_stride_a = _TYPE_SIZES[0]
sz_stride_b = _TYPE_SIZES[1]
sz_stride_c = _TYPE_SIZES[2]
sz_stride_d = _TYPE_SIZES[3]
sz_layout_sfa = _TYPE_SIZES[4]
sz_layout_sfb = _TYPE_SIZES[5]
sz_problem_shape = _TYPE_SIZES[6]
mnk = torch.zeros(ng * 3, dtype=torch.int32)
for i, (m, n, k) in enumerate(problem_sizes_list):
mnk[i*3+0] = m
mnk[i*3+1] = n
mnk[i*3+2] = k
h_ps = torch.zeros(ng * sz_problem_shape, dtype=torch.uint8)
h_sa = torch.zeros(ng * sz_stride_a, dtype=torch.uint8)
h_sb = torch.zeros(ng * sz_stride_b, dtype=torch.uint8)
h_sc = torch.zeros(ng * sz_stride_c, dtype=torch.uint8)
h_sd = torch.zeros(ng * sz_stride_d, dtype=torch.uint8)
h_lsfa = torch.zeros(ng * sz_layout_sfa, dtype=torch.uint8)
h_lsfb = torch.zeros(ng * sz_layout_sfb, dtype=torch.uint8)
meta_ret = _MODULE.fill_meta(mode, ng, mnk, h_ps, h_sa, h_sb, h_sc, h_sd, h_lsfa, h_lsfb)
if meta_ret != 0:
raise RuntimeError(f"fill_meta failed for mode={mode}, code={meta_ret}")
device = a_list[0].device
d_ps = h_ps.to(device)
d_sa = h_sa.to(device)
d_sb = h_sb.to(device)
d_sc = h_sc.to(device)
d_sd = h_sd.to(device)
d_lsfa = h_lsfa.to(device)
d_lsfb = h_lsfb.to(device)
ptr_a = torch.tensor([t.data_ptr() for t in a_list], dtype=torch.int64, device=device)
ptr_b = torch.tensor([t.data_ptr() for t in b_list], dtype=torch.int64, device=device)
ptr_sfa = torch.tensor([t.data_ptr() for t in sfa_list], dtype=torch.int64, device=device)
ptr_sfb = torch.tensor([t.data_ptr() for t in sfb_list], dtype=torch.int64, device=device)
ptr_c = torch.tensor([t.data_ptr() for t in out_list], dtype=torch.int64, device=device)
ptr_d = torch.tensor([t.data_ptr() for t in out_list], dtype=torch.int64, device=device)
return {
"d_ps": d_ps, "h_ps": h_ps,
"ptr_a": ptr_a, "d_sa": d_sa,
"ptr_b": ptr_b, "d_sb": d_sb,
"ptr_sfa": ptr_sfa, "d_lsfa": d_lsfa,
"ptr_sfb": ptr_sfb, "d_lsfb": d_lsfb,
"ptr_c": ptr_c, "d_sc": d_sc,
"ptr_d": ptr_d, "d_sd": d_sd,
}
_GROUPED_META_CACHE: Dict[Tuple, Dict[str, torch.Tensor]] = {}
_GROUPED_MODE_CACHE: Dict[Tuple[Tuple[int, int, int], ...], int] = {}
def _get_cutlass_metadata(problem_sizes_list, a_list, b_list, sfa_list, sfb_list, out_list, mode: int):
key = (
mode,
tuple(problem_sizes_list),
tuple(t.data_ptr() for t in a_list),
tuple(t.data_ptr() for t in b_list),
tuple(t.data_ptr() for t in sfa_list),
tuple(t.data_ptr() for t in sfb_list),
tuple(t.data_ptr() for t in out_list),
)
cached = _GROUPED_META_CACHE.get(key)
if cached is not None:
return cached
meta = _build_cutlass_metadata(problem_sizes_list, a_list, b_list, sfa_list, sfb_list, out_list, mode)
_GROUPED_META_CACHE[key] = meta
return meta
def _select_grouped_mode(ps_list, a_list, b_list, sfa_list, sfb_list, out_list):
key = tuple(ps_list)
cached = _GROUPED_MODE_CACHE.get(key)
if cached is not None:
return cached
try_mode_2sm = len(ps_list) >= 2 and min(k for _, _, k in ps_list) >= 2048
if try_mode_2sm:
meta2 = _get_cutlass_metadata(ps_list, a_list, b_list, sfa_list, sfb_list, out_list, 2)
ret2 = _MODULE.can_implement(
len(ps_list),
meta2["d_ps"],
meta2["ptr_a"], meta2["d_sa"],
meta2["ptr_b"], meta2["d_sb"],
meta2["ptr_sfa"], meta2["d_lsfa"],
meta2["ptr_sfb"], meta2["d_lsfb"],
meta2["ptr_c"], meta2["d_sc"],
meta2["ptr_d"], meta2["d_sd"],
_WORKSPACE,
meta2["h_ps"],
2,
)
if ret2 == 0:
print("[INFO] CUTLASS grouped mode selected: 2SM", file=sys.stderr, flush=True)
_GROUPED_MODE_CACHE[key] = 2
return 2
print(f"[INFO] CUTLASS grouped mode selected: 1SM (2SM unavailable, code={ret2})", file=sys.stderr, flush=True)
_GROUPED_MODE_CACHE[key] = 1
return 1
def _scaled_mm_cached_kernel(abc_tensors, sfasfb_tensors, problem_sizes):
result_tensors = []
for (a_ref, b_ref, c_ref), (sfa_ref, sfb_ref), (_, _, _, l) in zip(
abc_tensors,
sfasfb_tensors,
problem_sizes,
):
for l_idx in range(l):
a = a_ref[:, :, l_idx].view(torch.float4_e2m1fn_x2)
b = b_ref[:, :, l_idx].transpose(0, 1).view(torch.float4_e2m1fn_x2)
scale_a = _SCALE_CACHE.get_or_create(sfa_ref, l_idx, a.device)
scale_b = _SCALE_CACHE.get_or_create(sfb_ref, l_idx, a.device)
torch._scaled_mm(
a,
b,
scale_a,
scale_b,
bias=None,
out_dtype=torch.float16,
out=c_ref[:, :, l_idx],
)
result_tensors.append(c_ref)
return result_tensors
def ref_kernel(data: input_t) -> output_t:
abc_tensors, sfasfb_tensors, _, problem_sizes = _unpack_input(data)
return _scaled_mm_cached_kernel(abc_tensors, sfasfb_tensors, problem_sizes)
def custom_kernel(data: input_t) -> output_t:
abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = _unpack_input(data)
a_list = []
b_list = []
sfa_list = []
sfb_list = []
out_list = []
ps_list = []
use_reordered = sfasfb_reordered_tensors is not None
if not use_reordered:
print("[INFO] Grouped CUTLASS path using to_blocked scales", file=sys.stderr, flush=True)
if use_reordered:
loop_iter = zip(abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)
for (a_ref, b_ref, c_ref), (sfa_ref, sfb_ref), (sfa_reordered, sfb_reordered), (m, n, k, l) in loop_iter:
for l_idx in range(l):
a = a_ref.view(torch.uint8)[:, :, l_idx].contiguous()
b = b_ref.view(torch.uint8)[:, :, l_idx].contiguous()
sfa_dev = _REORDERED_SCALE_CACHE.get_or_create(sfa_reordered, l_idx, a.device)
sfb_dev = _REORDERED_SCALE_CACHE.get_or_create(sfb_reordered, l_idx, a.device)
a_list.append(a)
b_list.append(b)
sfa_list.append(sfa_dev)
sfb_list.append(sfb_dev)
out_list.append(c_ref[:, :, l_idx])
ps_list.append((int(m), int(n), int(k)))
else:
for (a_ref, b_ref, c_ref), (sfa_ref, sfb_ref), (m, n, k, l) in zip(
abc_tensors, sfasfb_tensors, problem_sizes
):
for l_idx in range(l):
a = a_ref.view(torch.uint8)[:, :, l_idx].contiguous()
b = b_ref.view(torch.uint8)[:, :, l_idx].contiguous()
sfa_dev = _SCALE_CACHE.get_or_create(sfa_ref, l_idx, a.device)
sfb_dev = _SCALE_CACHE.get_or_create(sfb_ref, l_idx, a.device)
a_list.append(a)
b_list.append(b)
sfa_list.append(sfa_dev)
sfb_list.append(sfb_dev)
out_list.append(c_ref[:, :, l_idx])
ps_list.append((int(m), int(n), int(k)))
mode = _select_grouped_mode(ps_list, a_list, b_list, sfa_list, sfb_list, out_list)
meta = _get_cutlass_metadata(ps_list, a_list, b_list, sfa_list, sfb_list, out_list, mode)
ret = _MODULE.run(
len(ps_list),
meta["d_ps"],
meta["ptr_a"], meta["d_sa"],
meta["ptr_b"], meta["d_sb"],
meta["ptr_sfa"], meta["d_lsfa"],
meta["ptr_sfb"], meta["d_lsfb"],
meta["ptr_c"], meta["d_sc"],
meta["ptr_d"], meta["d_sd"],
_WORKSPACE,
meta["h_ps"],
mode,
)
if ret != 0:
raise RuntimeError(f"CUTLASS grouped GEMM failed with code {ret} (mode={mode})")
return [c_ref for _, _, c_ref in abc_tensors]
def create_reordered_scale_factor_tensor(l, mn, k, ref_f8_tensor):
sf_k = ceil_div(k, sf_vec_size)
atom_m = (32, 4)
atom_k = 4
mma_shape = (
l,
ceil_div(mn, atom_m[0] * atom_m[1]),
ceil_div(sf_k, atom_k),
atom_m[0],
atom_m[1],
atom_k,
)
mma_permute_order = (3, 4, 1, 5, 2, 0)
rand_int_tensor = torch.randint(1, 3, mma_shape, dtype=torch.int8, device="cuda")
reordered_f8_tensor = rand_int_tensor.to(dtype=torch.float8_e4m3fn).permute(*mma_permute_order)
if ref_f8_tensor.device.type == "cpu":
ref_f8_tensor = ref_f8_tensor.cuda()
i_idx = torch.arange(mn, device="cuda")
j_idx = torch.arange(sf_k, device="cuda")
b_idx = torch.arange(l, device="cuda")
i_grid, j_grid, b_grid = torch.meshgrid(i_idx, j_idx, b_idx, indexing="ij")
mm = i_grid // (atom_m[0] * atom_m[1])
mm32 = i_grid % atom_m[0]
mm4 = (i_grid % 128) // atom_m[0]
kk = j_grid // atom_k
kk4 = j_grid % atom_k
reordered_f8_tensor[mm32, mm4, mm, kk4, kk, b_grid] = ref_f8_tensor[i_grid, j_grid, b_grid]
return reordered_f8_tensor
def _create_fp4_tensors(l, mn, k):
ref_i8 = torch.randint(255, size=(l, mn, k // 2), dtype=torch.uint8, device="cuda")
ref_i8 = ref_i8 & 0b1011_1011
return ref_i8.permute(1, 2, 0).view(torch.float4_e2m1fn_x2)
def generate_input(m: tuple, n: tuple, k: tuple, g: int, seed: int):
torch.manual_seed(seed)
abc_tensors = []
sfasfb_tensors = []
sfasfb_reordered_tensors = []
problem_sizes = []
l = 1
for group_idx in range(g):
mi = m[group_idx]
ni = n[group_idx]
ki = k[group_idx]
a_ref = _create_fp4_tensors(l, mi, ki)
b_ref = _create_fp4_tensors(l, ni, ki)
c_ref = torch.randn((l, mi, ni), dtype=torch.float16, device="cuda").permute(1, 2, 0)
sf_k = ceil_div(ki, sf_vec_size)
sfa_ref_cpu = torch.randint(1, 3, (l, mi, sf_k), dtype=torch.int8).to(
dtype=torch.float8_e4m3fn
).permute(1, 2, 0)
sfb_ref_cpu = torch.randint(1, 3, (l, ni, sf_k), dtype=torch.int8).to(
dtype=torch.float8_e4m3fn
).permute(1, 2, 0)
sfa_reordered = create_reordered_scale_factor_tensor(l, mi, ki, sfa_ref_cpu)
sfb_reordered = create_reordered_scale_factor_tensor(l, ni, ki, sfb_ref_cpu)
abc_tensors.append((a_ref, b_ref, c_ref))
sfasfb_tensors.append((sfa_ref_cpu, sfb_ref_cpu))
sfasfb_reordered_tensors.append((sfa_reordered, sfb_reordered))
problem_sizes.append((mi, ni, ki, l))
return (abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)
check_implementation = make_match_reference(ref_kernel, rtol=1e-03, atol=1e-03)
scrolls · 979 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