Skip to content
KernelIndex
Search⌘K

submission 479955

J · python · License unknown

Use it

Vendorable · source mirrored · license unknownView source →

No package. Vendor the mirrored source: 1116 lines, June 9 Researcher Reciprocity License v1.0.

submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-group-gemm-479955?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
42.5µs
#46 of 145
2026-02-06

Reported · How evidence levels are derived →

Source and license

sourceavailable
revision digestsha256:b7a52a30bc7a9cfe98dd9624bfeda4b4eba345c72a2b6006427996cc4d6a7dee
license declaredunknown
license concludedunknown
authorsJ
imported2026-08-15

Techniques

Extracted from the mirrored source by pattern, never inferred. Each row cites its line.

clusterusing ClusterShape = Shape<int32_t,int32_t,_1>;
fp4using ElementInput = cutlass::float_e2m1_t;
fused-epilogueusing EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;
warp-specializationusing EpilogueSchedule = cutlass::epilogue::PtrArrayTmaWarpSpecialized1Sm;

Kernel source

submission.py1116 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()
    );
}

struct RegisteredPlan {
    int num_groups;
    int mode;
    at::Tensor d_ps, d_ptr_A, d_sa, d_ptr_B, d_sb;
    at::Tensor d_ptr_SFA, d_lsfa, d_ptr_SFB, d_lsfb;
    at::Tensor d_ptr_C, d_sc, d_ptr_D, d_sd;
    at::Tensor workspace;
    at::Tensor h_ps;
};

static std::vector<RegisteredPlan> g_plans;

int register_plan(
    int num_groups,
    at::Tensor d_ps,
    at::Tensor d_ptr_A, at::Tensor d_sa,
    at::Tensor d_ptr_B, at::Tensor d_sb,
    at::Tensor d_ptr_SFA, at::Tensor d_lsfa,
    at::Tensor d_ptr_SFB, at::Tensor d_lsfb,
    at::Tensor d_ptr_C, at::Tensor d_sc,
    at::Tensor d_ptr_D, at::Tensor d_sd,
    at::Tensor workspace,
    at::Tensor h_ps,
    int mode
) {
    int ret = run_cutlass_grouped_gemm_mode(
        num_groups,
        d_ps.data_ptr(),
        d_ptr_A.data_ptr(), d_sa.data_ptr(),
        d_ptr_B.data_ptr(), d_sb.data_ptr(),
        d_ptr_SFA.data_ptr(), d_lsfa.data_ptr(),
        d_ptr_SFB.data_ptr(), d_lsfb.data_ptr(),
        d_ptr_C.data_ptr(), d_sc.data_ptr(),
        d_ptr_D.data_ptr(), d_sd.data_ptr(),
        workspace.data_ptr(), workspace.nbytes(),
        h_ps.data_ptr(),
        mode,
        false
    );
    if (ret != 0) return ret;

    ret = run_cutlass_grouped_gemm_mode(
        num_groups,
        d_ps.data_ptr(),
        d_ptr_A.data_ptr(), d_sa.data_ptr(),
        d_ptr_B.data_ptr(), d_sb.data_ptr(),
        d_ptr_SFA.data_ptr(), d_lsfa.data_ptr(),
        d_ptr_SFB.data_ptr(), d_lsfb.data_ptr(),
        d_ptr_C.data_ptr(), d_sc.data_ptr(),
        d_ptr_D.data_ptr(), d_sd.data_ptr(),
        workspace.data_ptr(), workspace.nbytes(),
        h_ps.data_ptr(),
        mode,
        true
    );
    if (ret != 0) return ret;

    RegisteredPlan plan;
    plan.num_groups = num_groups;
    plan.mode = mode;
    plan.d_ps = d_ps;
    plan.d_ptr_A = d_ptr_A; plan.d_sa = d_sa;
    plan.d_ptr_B = d_ptr_B; plan.d_sb = d_sb;
    plan.d_ptr_SFA = d_ptr_SFA; plan.d_lsfa = d_lsfa;
    plan.d_ptr_SFB = d_ptr_SFB; plan.d_lsfb = d_lsfb;
    plan.d_ptr_C = d_ptr_C; plan.d_sc = d_sc;
    plan.d_ptr_D = d_ptr_D; plan.d_sd = d_sd;
    plan.workspace = workspace;
    plan.h_ps = h_ps;
    int idx = g_plans.size();
    g_plans.push_back(plan);
    return idx;
}

int run_plan(int plan_idx) {
    if (plan_idx < 0 || plan_idx >= (int)g_plans.size()) return -100;
    auto& p = g_plans[plan_idx];
    return run_cutlass_grouped_gemm_mode(
        p.num_groups,
        p.d_ps.data_ptr(),
        p.d_ptr_A.data_ptr(), p.d_sa.data_ptr(),
        p.d_ptr_B.data_ptr(), p.d_sb.data_ptr(),
        p.d_ptr_SFA.data_ptr(), p.d_lsfa.data_ptr(),
        p.d_ptr_SFB.data_ptr(), p.d_lsfb.data_ptr(),
        p.d_ptr_C.data_ptr(), p.d_sc.data_ptr(),
        p.d_ptr_D.data_ptr(), p.d_sd.data_ptr(),
        p.workspace.data_ptr(), p.workspace.nbytes(),
        p.h_ps.data_ptr(),
        p.mode,
        true
    );
}

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("register_plan", &register_plan, "Register and warmup a plan");
    m.def("run_plan", &run_plan, "Run a registered plan");
    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] = {}
_GROUPED_PLAN_CACHE: Dict[Tuple, Dict] = {}


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) >= 256
    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 _build_grouped_plan(abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes):
    a_list = []
    b_list = []
    sfa_list = []
    sfb_list = []
    out_list = []
    ps_list = []

    use_reordered = sfasfb_reordered_tensors is not None

    if use_reordered:
        loop_iter = zip(abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)
        for (a_ref, b_ref, c_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)

    plan_id = _MODULE.register_plan(
        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 plan_id < 0:
        raise RuntimeError(f"register_plan failed: {plan_id}")

    return {"plan_id": plan_id, "mode": mode}


def _get_grouped_plan(abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes):
    use_reordered = sfasfb_reordered_tensors is not None
    key_parts = [use_reordered]
    if use_reordered:
        for (a_ref, b_ref, c_ref), (sfa_reordered, sfb_reordered), (m, n, k, l) in zip(
            abc_tensors, sfasfb_reordered_tensors, problem_sizes
        ):
            key_parts.append((
                id(a_ref), a_ref._version,
                id(b_ref), b_ref._version,
                id(c_ref),
                id(sfa_reordered), sfa_reordered._version,
                id(sfb_reordered), sfb_reordered._version,
                int(m), int(n), int(k), int(l),
            ))
    else:
        for (a_ref, b_ref, c_ref), (sfa_ref, sfb_ref), (m, n, k, l) in zip(
            abc_tensors, sfasfb_tensors, problem_sizes
        ):
            key_parts.append((
                id(a_ref), a_ref._version,
                id(b_ref), b_ref._version,
                id(c_ref),
                id(sfa_ref), sfa_ref._version,
                id(sfb_ref), sfb_ref._version,
                int(m), int(n), int(k), int(l),
            ))
    key = tuple(key_parts)
    cached = _GROUPED_PLAN_CACHE.get(key)
    if cached is not None:
        return cached
    plan = _build_grouped_plan(abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)
    _GROUPED_PLAN_CACHE[key] = plan
    return plan


def custom_kernel(data: input_t) -> output_t:
    abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = _unpack_input(data)

    plan = _get_grouped_plan(abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)

    ret = _MODULE.run_plan(plan["plan_id"])
    if ret != 0:
        raise RuntimeError(f"CUTLASS grouped GEMM failed with code {ret} (mode={plan['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 · 1116 lines total

Source code from GPU Mode and the KernelBot dataset · June 9 Researcher Reciprocity License v1.0

Changes from previous submission

Against this author's previous submission submission 476965.

⋯ 601 unchanged lines
);
}
+ struct RegisteredPlan {
+ int num_groups;
+ int mode;
+ at::Tensor d_ps, d_ptr_A, d_sa, d_ptr_B, d_sb;
+ at::Tensor d_ptr_SFA, d_lsfa, d_ptr_SFB, d_lsfb;
+ at::Tensor d_ptr_C, d_sc, d_ptr_D, d_sd;
+ at::Tensor workspace;
+ at::Tensor h_ps;
+ };
+
+ static std::vector<RegisteredPlan> g_plans;
+
+ int register_plan(
+ int num_groups,
+ at::Tensor d_ps,
+ at::Tensor d_ptr_A, at::Tensor d_sa,
+ at::Tensor d_ptr_B, at::Tensor d_sb,
+ at::Tensor d_ptr_SFA, at::Tensor d_lsfa,
+ at::Tensor d_ptr_SFB, at::Tensor d_lsfb,
+ at::Tensor d_ptr_C, at::Tensor d_sc,
+ at::Tensor d_ptr_D, at::Tensor d_sd,
+ at::Tensor workspace,
+ at::Tensor h_ps,
+ int mode
+ ) {
+ int ret = run_cutlass_grouped_gemm_mode(
+ num_groups,
+ d_ps.data_ptr(),
+ d_ptr_A.data_ptr(), d_sa.data_ptr(),
+ d_ptr_B.data_ptr(), d_sb.data_ptr(),
+ d_ptr_SFA.data_ptr(), d_lsfa.data_ptr(),
+ d_ptr_SFB.data_ptr(), d_lsfb.data_ptr(),
+ d_ptr_C.data_ptr(), d_sc.data_ptr(),
+ d_ptr_D.data_ptr(), d_sd.data_ptr(),
+ workspace.data_ptr(), workspace.nbytes(),
+ h_ps.data_ptr(),
+ mode,
+ false
+ );
+ if (ret != 0) return ret;
+
+ ret = run_cutlass_grouped_gemm_mode(
+ num_groups,
+ d_ps.data_ptr(),
+ d_ptr_A.data_ptr(), d_sa.data_ptr(),
+ d_ptr_B.data_ptr(), d_sb.data_ptr(),
+ d_ptr_SFA.data_ptr(), d_lsfa.data_ptr(),
+ d_ptr_SFB.data_ptr(), d_lsfb.data_ptr(),
+ d_ptr_C.data_ptr(), d_sc.data_ptr(),
+ d_ptr_D.data_ptr(), d_sd.data_ptr(),
+ workspace.data_ptr(), workspace.nbytes(),
+ h_ps.data_ptr(),
+ mode,
+ true
+ );
+ if (ret != 0) return ret;
+
+ RegisteredPlan plan;
+ plan.num_groups = num_groups;
+ plan.mode = mode;
+ plan.d_ps = d_ps;
+ plan.d_ptr_A = d_ptr_A; plan.d_sa = d_sa;
+ plan.d_ptr_B = d_ptr_B; plan.d_sb = d_sb;
+ plan.d_ptr_SFA = d_ptr_SFA; plan.d_lsfa = d_lsfa;
+ plan.d_ptr_SFB = d_ptr_SFB; plan.d_lsfb = d_lsfb;
+ plan.d_ptr_C = d_ptr_C; plan.d_sc = d_sc;
+ plan.d_ptr_D = d_ptr_D; plan.d_sd = d_sd;
+ plan.workspace = workspace;
+ plan.h_ps = h_ps;
+ int idx = g_plans.size();
+ g_plans.push_back(plan);
+ return idx;
+ }
+
+ int run_plan(int plan_idx) {
+ if (plan_idx < 0 || plan_idx >= (int)g_plans.size()) return -100;
+ auto& p = g_plans[plan_idx];
+ return run_cutlass_grouped_gemm_mode(
+ p.num_groups,
+ p.d_ps.data_ptr(),
+ p.d_ptr_A.data_ptr(), p.d_sa.data_ptr(),
+ p.d_ptr_B.data_ptr(), p.d_sb.data_ptr(),
+ p.d_ptr_SFA.data_ptr(), p.d_lsfa.data_ptr(),
+ p.d_ptr_SFB.data_ptr(), p.d_lsfb.data_ptr(),
+ p.d_ptr_C.data_ptr(), p.d_sc.data_ptr(),
+ p.d_ptr_D.data_ptr(), p.d_sd.data_ptr(),
+ p.workspace.data_ptr(), p.workspace.nbytes(),
+ p.h_ps.data_ptr(),
+ p.mode,
+ true
+ );
+ }
+
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("register_plan", &register_plan, "Register and warmup a plan");
+ m.def("run_plan", &run_plan, "Run a registered plan");
m.def("can_implement", &can_implement_grouped, "CUTLASS grouped FP4 GEMM can_implement");
m.def("fill_meta", &fill_meta, "Fill metadata arrays");
}
⋯ 136 unchanged lines
_GROUPED_META_CACHE: Dict[Tuple, Dict[str, torch.Tensor]] = {}
_GROUPED_MODE_CACHE: Dict[Tuple[Tuple[int, int, int], ...], int] = {}
+ _GROUPED_PLAN_CACHE: Dict[Tuple, Dict] = {}
def _get_cutlass_metadata(problem_sizes_list, a_list, b_list, sfa_list, sfb_list, out_list, mode: int):
⋯ 20 unchanged lines
if cached is not None:
return cached
- try_mode_2sm = len(ps_list) >= 2 and min(k for _, _, k in ps_list) >= 2048
+ try_mode_2sm = len(ps_list) >= 2 and min(k for _, _, k in ps_list) >= 256
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(
⋯ 49 unchanged lines
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)
-
+ def _build_grouped_plan(abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes):
a_list = []
b_list = []
sfa_list = []
⋯ 2 unchanged lines
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 (a_ref, b_ref, c_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)
⋯ 9 unchanged lines
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)
⋯ 4 unchanged lines
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(
+ plan_id = _MODULE.register_plan(
len(ps_list),
meta["d_ps"],
meta["ptr_a"], meta["d_sa"],
⋯ 6 unchanged lines
meta["h_ps"],
mode,
)
+ if plan_id < 0:
+ raise RuntimeError(f"register_plan failed: {plan_id}")
+ return {"plan_id": plan_id, "mode": mode}
+
+
+ def _get_grouped_plan(abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes):
+ use_reordered = sfasfb_reordered_tensors is not None
+ key_parts = [use_reordered]
+ if use_reordered:
+ for (a_ref, b_ref, c_ref), (sfa_reordered, sfb_reordered), (m, n, k, l) in zip(
+ abc_tensors, sfasfb_reordered_tensors, problem_sizes
+ ):
+ key_parts.append((
+ id(a_ref), a_ref._version,
+ id(b_ref), b_ref._version,
+ id(c_ref),
+ id(sfa_reordered), sfa_reordered._version,
+ id(sfb_reordered), sfb_reordered._version,
+ int(m), int(n), int(k), int(l),
+ ))
+ else:
+ for (a_ref, b_ref, c_ref), (sfa_ref, sfb_ref), (m, n, k, l) in zip(
+ abc_tensors, sfasfb_tensors, problem_sizes
+ ):
+ key_parts.append((
+ id(a_ref), a_ref._version,
+ id(b_ref), b_ref._version,
+ id(c_ref),
+ id(sfa_ref), sfa_ref._version,
+ id(sfb_ref), sfb_ref._version,
+ int(m), int(n), int(k), int(l),
+ ))
+ key = tuple(key_parts)
+ cached = _GROUPED_PLAN_CACHE.get(key)
+ if cached is not None:
+ return cached
+ plan = _build_grouped_plan(abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)
+ _GROUPED_PLAN_CACHE[key] = plan
+ return plan
+
+
+ def custom_kernel(data: input_t) -> output_t:
+ abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes = _unpack_input(data)
+
+ plan = _get_grouped_plan(abc_tensors, sfasfb_tensors, sfasfb_reordered_tensors, problem_sizes)
+
+ ret = _MODULE.run_plan(plan["plan_id"])
if ret != 0:
- raise RuntimeError(f"CUTLASS grouped GEMM failed with code {ret} (mode={mode})")
+ raise RuntimeError(f"CUTLASS grouped GEMM failed with code {ret} (mode={plan['mode']})")
return [c_ref for _, _, c_ref in abc_tensors]
scrolls · 228 diff lines total

Best evidence level for this revision: reported

JSON