submission 497547
Darshan · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 355 lines, June 9 Researcher Reciprocity License v1.0.
v6.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-497547?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:9c443a784b514a1986ebda0904b5b9deff05efeca270e6cbbbcb902d5e7a1d52
license declaredunknown
license concludedunknown
authorsDarshan
imported2026-08-26
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 ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;fused-epilogue
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;warp-specialization
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;Kernel source
v6.py355 lines
import torch
import os
from torch.utils.cpp_extension import load_inline
input_t = tuple
output_t = torch.Tensor
CUDA_SRC = r"""
#include "cute/tensor.hpp"
#include "cutlass/cutlass.h"
#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/dispatch_policy.hpp"
#include "cutlass/gemm/group_array_problem_shape.hpp"
#include "cutlass/gemm/kernel/gemm_universal.hpp"
#include "cutlass/tensor_ref.h"
#include "cutlass/util/packed_stride.hpp"
#include <cuda_runtime.h>
#include <ATen/core/Tensor.h>
#include <torch/library.h>
#include <torch/types.h>
using namespace cute;
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
using ProblemShape = cutlass::gemm::GroupProblemShape<Shape<int, int, int>>;
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutA = cutlass::layout::RowMajor;
constexpr int AlignmentA = 32;
using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
using LayoutB = cutlass::layout::ColumnMajor;
constexpr int AlignmentB = 32;
using ElementC = cutlass::half_t;
using ElementD = cutlass::half_t;
using LayoutC = cutlass::layout::RowMajor;
using LayoutD = 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 MmaTileShape = Shape<_128, _256, _256>;
using ClusterShape = Shape<int32_t, int32_t, _1>;
using KernelSchedule = cutlass::gemm::KernelPtrArrayTmaWarpSpecialized1SmNvf4Sm100;
using EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;
using CollectiveEpilogue =
typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag, OperatorClass,
MmaTileShape, ClusterShape,
Shape<_128, _64>,
ElementAccumulator, ElementAccumulator,
ElementC, LayoutC *, AlignmentC,
ElementD, LayoutD *, AlignmentD,
EpilogueSchedule>::CollectiveOp;
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
ArchTag, OperatorClass,
ElementA, LayoutA *, AlignmentA,
ElementB, LayoutB *, AlignmentB,
ElementAccumulator,
MmaTileShape, ClusterShape,
cutlass::gemm::collective::StageCountAutoCarveout<
static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
KernelSchedule
>::CollectiveOp;
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
ProblemShape,
CollectiveMainloop,
CollectiveEpilogue>;
using Gemm = cutlass::gemm::device::GemmUniversalAdapter<GemmKernel>;
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 LayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFA;
using LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::InternalLayoutSFB;
using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
using ElementSF = typename Gemm::GemmKernel::ElementSF;
// Persistent state cached across calls
static int g_sm_count = -1;
static char* g_pinned_host = nullptr;
static size_t g_pinned_size = 0;
static at::Tensor g_device_buf;
static char* g_device_ptr = nullptr;
static size_t g_device_size = 0;
static at::Tensor g_workspace_buf;
static void* g_workspace_ptr = nullptr;
static size_t g_workspace_size = 0;
void nvfp4_grouped_gemm(
at::TensorList a,
at::TensorList b,
at::TensorList sfa,
at::TensorList sfb,
at::TensorList d,
at::IntArrayRef ms,
at::IntArrayRef ns,
at::IntArrayRef ks)
{
int num_groups = static_cast<int>(a.size());
TORCH_CHECK(num_groups > 0, "Need at least one group");
if (g_sm_count < 0) {
g_sm_count = cutlass::KernelHardwareInfo::query_device_multiprocessor_count(0);
}
using UnderlyingProblemShape = typename ProblemShape::UnderlyingProblemShape;
constexpr size_t ALIGN = 16;
auto align_up = [](size_t x, size_t a) { return (x + a - 1) & ~(a - 1); };
size_t off = 0;
size_t off_ps = off; off += align_up(num_groups * sizeof(UnderlyingProblemShape), ALIGN);
size_t off_pA = off; off += align_up(num_groups * sizeof(const void*), ALIGN);
size_t off_pB = off; off += align_up(num_groups * sizeof(const void*), ALIGN);
size_t off_pSFA = off; off += align_up(num_groups * sizeof(const void*), ALIGN);
size_t off_pSFB = off; off += align_up(num_groups * sizeof(const void*), ALIGN);
size_t off_pC = off; off += align_up(num_groups * sizeof(const void*), ALIGN);
size_t off_pD = off; off += align_up(num_groups * sizeof(void*), ALIGN);
size_t off_sA = off; off += align_up(num_groups * sizeof(StrideA), ALIGN);
size_t off_sB = off; off += align_up(num_groups * sizeof(StrideB), ALIGN);
size_t off_sC = off; off += align_up(num_groups * sizeof(StrideC), ALIGN);
size_t off_sD = off; off += align_up(num_groups * sizeof(StrideD), ALIGN);
size_t off_lSFA = off; off += align_up(num_groups * sizeof(LayoutSFA), ALIGN);
size_t off_lSFB = off; off += align_up(num_groups * sizeof(LayoutSFB), ALIGN);
size_t total = off;
// Persistent pinned host buffer
if (total > g_pinned_size) {
if (g_pinned_host) cudaFreeHost(g_pinned_host);
size_t alloc = std::max(total, (size_t)65536);
cudaHostAlloc(&g_pinned_host, alloc, cudaHostAllocDefault);
g_pinned_size = alloc;
}
char* h = g_pinned_host;
memset(h, 0, total);
auto* ps = reinterpret_cast<UnderlyingProblemShape*>(h + off_ps);
auto* pA = reinterpret_cast<const void**>(h + off_pA);
auto* pB = reinterpret_cast<const void**>(h + off_pB);
auto* pSFA = reinterpret_cast<const void**>(h + off_pSFA);
auto* pSFB = reinterpret_cast<const void**>(h + off_pSFB);
auto* pC = reinterpret_cast<const void**>(h + off_pC);
auto* pD = reinterpret_cast<void**>(h + off_pD);
auto* sA = reinterpret_cast<StrideA*>(h + off_sA);
auto* sB = reinterpret_cast<StrideB*>(h + off_sB);
auto* sC = reinterpret_cast<StrideC*>(h + off_sC);
auto* sD = reinterpret_cast<StrideD*>(h + off_sD);
auto* lSFA = reinterpret_cast<LayoutSFA*>(h + off_lSFA);
auto* lSFB = reinterpret_cast<LayoutSFB*>(h + off_lSFB);
for (int i = 0; i < num_groups; i++) {
int M = static_cast<int>(ms[i]);
int N = static_cast<int>(ns[i]);
int K = static_cast<int>(ks[i]);
ps[i] = {M, N, K};
pA[i] = a[i].data_ptr();
pB[i] = b[i].data_ptr();
pSFA[i] = sfa[i].data_ptr();
pSFB[i] = sfb[i].data_ptr();
pC[i] = nullptr;
pD[i] = d[i].data_ptr();
sA[i] = cutlass::make_cute_packed_stride(StrideA{}, {M, K, 1});
sB[i] = cutlass::make_cute_packed_stride(StrideB{}, {N, K, 1});
sC[i] = cutlass::make_cute_packed_stride(StrideC{}, {M, N, 1});
sD[i] = cutlass::make_cute_packed_stride(StrideD{}, {M, N, 1});
lSFA[i] = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(make_shape(M, N, K, 1));
lSFB[i] = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(make_shape(M, N, K, 1));
}
// Persistent device buffer
if (total > g_device_size) {
size_t alloc = std::max(total, (size_t)65536);
g_device_buf = at::empty({(int64_t)alloc},
at::TensorOptions().dtype(at::kByte).device(a[0].device()));
g_device_ptr = (char*)g_device_buf.data_ptr();
g_device_size = alloc;
}
cudaMemcpy(g_device_ptr, h, total, cudaMemcpyHostToDevice);
// Cluster 1x1 — optimal for small M values (40-384)
cutlass::KernelHardwareInfo hw_info;
hw_info.device_id = 0;
hw_info.sm_count = g_sm_count;
hw_info.cluster_shape = dim3(1, 1, 1);
hw_info.cluster_shape_fallback = dim3(1, 1, 1);
typename Gemm::Arguments arguments;
decltype(arguments.epilogue.thread) fusion_args;
fusion_args.alpha = 1.0f;
fusion_args.beta = 0.0f;
fusion_args.alpha_ptr = nullptr;
fusion_args.beta_ptr = nullptr;
fusion_args.alpha_ptr_array = nullptr;
fusion_args.beta_ptr_array = nullptr;
fusion_args.dAlpha = {_0{}, _0{}, 0};
fusion_args.dBeta = {_0{}, _0{}, 0};
typename Gemm::GemmKernel::TileSchedulerArguments scheduler;
arguments = typename Gemm::Arguments{
cutlass::gemm::GemmUniversalMode::kGrouped,
{num_groups,
reinterpret_cast<UnderlyingProblemShape*>(g_device_ptr + off_ps),
ps},
{reinterpret_cast<const typename Gemm::ElementA **>(g_device_ptr + off_pA),
reinterpret_cast<StrideA*>(g_device_ptr + off_sA),
reinterpret_cast<const typename Gemm::ElementB **>(g_device_ptr + off_pB),
reinterpret_cast<StrideB*>(g_device_ptr + off_sB),
reinterpret_cast<const ElementSF **>(g_device_ptr + off_pSFA),
reinterpret_cast<LayoutSFA*>(g_device_ptr + off_lSFA),
reinterpret_cast<const ElementSF **>(g_device_ptr + off_pSFB),
reinterpret_cast<LayoutSFB*>(g_device_ptr + off_lSFB)},
{fusion_args,
reinterpret_cast<const ElementC **>(g_device_ptr + off_pC),
reinterpret_cast<StrideC*>(g_device_ptr + off_sC),
reinterpret_cast<ElementD **>(g_device_ptr + off_pD),
reinterpret_cast<StrideD*>(g_device_ptr + off_sD)},
hw_info, scheduler
};
Gemm gemm;
size_t workspace_size = Gemm::get_workspace_size(arguments);
void* workspace = nullptr;
if (workspace_size > 0) {
if (workspace_size > g_workspace_size) {
g_workspace_buf = at::empty({(int64_t)workspace_size},
at::TensorOptions().dtype(at::kByte).device(a[0].device()));
g_workspace_ptr = g_workspace_buf.data_ptr();
g_workspace_size = workspace_size;
}
workspace = g_workspace_ptr;
}
auto status = gemm.initialize(arguments, workspace);
TORCH_CHECK(status == cutlass::Status::kSuccess,
"CUTLASS grouped GEMM initialize failed");
status = gemm.run();
TORCH_CHECK(status == cutlass::Status::kSuccess,
"CUTLASS grouped GEMM run failed");
}
#else
void nvfp4_grouped_gemm(
at::TensorList, at::TensorList,
at::TensorList, at::TensorList,
at::TensorList,
at::IntArrayRef, at::IntArrayRef, at::IntArrayRef) {
TORCH_CHECK(false, "SM100 not supported");
}
#endif
TORCH_LIBRARY(nvfp4_v6, m) {
m.def("nvfp4_grouped_gemm(Tensor[] a, Tensor[] b, Tensor[] sfa, Tensor[] sfb, Tensor[] d, int[] ms, int[] ns, int[] ks) -> ()");
m.impl("nvfp4_grouped_gemm", &nvfp4_grouped_gemm);
}
"""
cutlass_path = os.environ.get("CUTLASS_PATH", "/mnt/Code/cutlass")
cuda_include = os.environ.get("CUDA_INCLUDE_DIR", "/usr/local/cuda/include")
load_inline(
"nvfp4_grouped_gemm_v6",
cpp_sources="",
cuda_sources=CUDA_SRC,
verbose=True,
is_python_module=False,
no_implicit_headers=True,
extra_include_paths=[
f"{cutlass_path}/include",
f"{cutlass_path}/tools/util/include",
cuda_include,
],
extra_cuda_cflags=[
"-std=c++17",
"-gencode=arch=compute_100a,code=sm_100a",
"-DCUTLASS_ARCH_MMA_SM100_SUPPORTED=1",
"-O3", "--use_fast_math",
"--ftz=true", "--prec-div=false", "--prec-sqrt=false",
"--expt-relaxed-constexpr",
"--relocatable-device-code=false",
"-lineinfo",
"-Xptxas=-v",
],
extra_ldflags=["-lcuda"],
)
grouped_gemm = torch.ops.nvfp4_v6.nvfp4_grouped_gemm
def custom_kernel(data: input_t) -> output_t:
abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = data
a_list = []
b_list = []
sfa_list = []
sfb_list = []
d_list = []
ms_list = []
ns_list = []
ks_list = []
need_copyback = []
for (a_ref, b_ref, c_ref), (sfa_r, sfb_r), (m, n, k, l) in zip(
abc_tensors, sfasfb_reordered_tensors, problem_sizes
):
for l_idx in range(l):
a_slice = a_ref[:, :, l_idx]
b_slice = b_ref[:, :, l_idx]
d_slice = c_ref[:, :, l_idx]
if a_slice.is_contiguous() and d_slice.is_contiguous():
a_list.append(a_slice)
b_list.append(b_slice)
d_list.append(d_slice)
else:
a_list.append(a_slice.contiguous())
b_list.append(b_slice.contiguous())
d_tmp = torch.empty((m, n), dtype=torch.float16, device=c_ref.device)
d_list.append(d_tmp)
need_copyback.append((d_tmp, c_ref, l_idx))
sfa_list.append(sfa_r)
sfb_list.append(sfb_r)
ms_list.append(m)
ns_list.append(n)
ks_list.append(k)
grouped_gemm(a_list, b_list, sfa_list, sfb_list, d_list, ms_list, ns_list, ks_list)
for d_tmp, c_ref, l_idx in need_copyback:
c_ref[:, :, l_idx].copy_(d_tmp)
return [c for (_, _, c) in abc_tensors]
scrolls · 355 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