Skip to content
KernelIndex
Search⌘K

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
NVFP4 group GEMMsuite of 4 cases
NVIDIA B200
66.7µs
#78 of 145
2026-02-18

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.

clusterusing ClusterShape = Shape<int32_t, int32_t, _1>;
fp4using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>;
fused-epilogueusing EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;
warp-specializationusing 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