submission 121654
jihun_lee__ · python · License unknown
Use it
Vendorable · source mirrored · license unknownView source →
No package. Vendor the mirrored source: 1706 lines, June 9 Researcher Reciprocity License v1.0.
submission.py
curl "https://kernelindex.com/api/v1/implementations/kernelbot-nvfp4-gemm-121654?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:97d6afcf8ec1a37e53ea008ad7b3da2338f049fc7e8998e876b66090ec2370c9
license declaredunknown
license concludedunknown
authorsjihun_lee__
imported2026-08-26
Techniques
Extracted from the mirrored source by pattern, never inferred. Each row cites its line.
cluster
using ClusterShape = cutlass::gemm::GemmShape<fp4
cuda_source = """// GENERATED CUDA CODE FOR NVFP4 BLOCKSCALED GEMVfused-epilogue
using CollectiveEpilogue = typename GemmKernel::CollectiveEpilogue;mbarrier
alignas(16) cutlass::arch::ClusterBarrier tmem_dealloc;persistent-kernel
static constexpr bool IsSchedDynamicPersistent = TileScheduler::IsDynamicPersistent;shared-memory
static int maximum_active_blocks(int /* smem_capacity */ = -1) {split-k
static int constexpr kSplitKAlignment = cute::max(tma
collective_mainloop.prefetch_tma_descriptors();warp-specialization
bool is_epi_load_needed = collective_epilogue.is_producer_load_needed();Kernel source
submission.py1706 lines
import torch
from torch.utils.cpp_extension import load_inline
from task import input_t, output_t
cuda_source = """// GENERATED CUDA CODE FOR NVFP4 BLOCKSCALED GEMV
// Generated: 2025-12-04 18:46:43
#include <torch/extension.h>
#include <cute/tensor.hpp>
#include <cute/algorithm/copy.hpp>
#include <cute/algorithm/gemm.hpp>
#include <cutlass/cutlass.h>
#include <cutlass/numeric_conversion.h>
#include <cutlass/arch/memory_sm80.h>
#include <cutlass/numeric_types.h>
// ========== gemm_universal_adapter.h ==========
// common
#include <cutlass/cutlass.h>
#include <cutlass/device_kernel.h>
#include <cutlass/gemm/gemm.h>
#include <cutlass/detail/layout.hpp>
#include <cutlass/detail/mma.hpp>
#include <cutlass/cuda_host_adapter.hpp>
#include <cutlass/kernel_launch.h>
#if !defined(__CUDACC_RTC__)
#include <cutlass/cluster_launch.hpp>
#include <cutlass/trace.h>
#endif // !defined(__CUDACC_RTC__)
// 3.x
#include <cutlass/gemm/kernel/gemm_universal.hpp>
#include <cutlass/gemm/device/gemm_universal_adapter.h>
////////////////////////////////////////////////////////////////////////////////
namespace gpu_mode {
namespace detail {
// Work-around for some DispatchPolicy types not having a Stages member.
// In that case, the Stages value is 0. Most code should static_assert
// that the number of stages is valid.
// Whether DispatchPolicy::Stages is valid.
// It should also be convertible to int, but if not, that will show up
// as a build error when GemmUniversalAdapter attempts to assign it to kStages.
template <class DispatchPolicy, class Enable = void>
struct has_Stages : cute::false_type {};
template <class DispatchPolicy>
struct has_Stages<DispatchPolicy, cute::void_t<decltype(DispatchPolicy::Stages)>> : cute::true_type {};
template<class DispatchPolicy>
constexpr int stages_member(DispatchPolicy) {
if constexpr (has_Stages<DispatchPolicy>::value) {
return DispatchPolicy::Stages;
}
else {
return 0;
}
}
template <class GemmKernel, class = void>
struct IsDistGemmKernel : cute::false_type { };
template <typename GemmKernel>
struct IsDistGemmKernel<GemmKernel, cute::void_t<typename GemmKernel::TP>>
: cute::true_type { };
} // namespace detail
template <class GemmKernel_>
class GemmUniversalAdapterCustom : cutlass::gemm::device::GemmUniversalAdapter<GemmKernel_>
{
public:
using GemmKernel = cutlass::GetUnderlyingKernel_t<GemmKernel_>;
using TileShape = typename GemmKernel::TileShape;
using ElementA = typename GemmKernel::ElementA;
using ElementB = typename GemmKernel::ElementB;
using ElementC = typename GemmKernel::ElementC;
using ElementD = typename GemmKernel::ElementD;
using ElementAccumulator = typename GemmKernel::ElementAccumulator;
using DispatchPolicy = typename GemmKernel::DispatchPolicy;
using CollectiveMainloop = typename GemmKernel::CollectiveMainloop;
using CollectiveEpilogue = typename GemmKernel::CollectiveEpilogue;
// Map back to 2.x type as best as possible
using LayoutA = cutlass::gemm::detail::StrideToLayoutTagA_t<typename GemmKernel::StrideA>;
using LayoutB = cutlass::gemm::detail::StrideToLayoutTagB_t<typename GemmKernel::StrideB>;
using LayoutC = cutlass::gemm::detail::StrideToLayoutTagC_t<typename GemmKernel::StrideC>;
using LayoutD = cutlass::gemm::detail::StrideToLayoutTagC_t<typename GemmKernel::StrideD>;
static bool const kEnableCudaHostAdapter = CUTLASS_ENABLE_CUDA_HOST_ADAPTER;
static cutlass::ComplexTransform const kTransformA = cute::is_same_v<typename GemmKernel::CollectiveMainloop::TransformA, cute::conjugate> ?
cutlass::ComplexTransform::kConjugate : cutlass::ComplexTransform::kNone;
static cutlass::ComplexTransform const kTransformB = cute::is_same_v<typename GemmKernel::CollectiveMainloop::TransformB, cute::conjugate> ?
cutlass::ComplexTransform::kConjugate : cutlass::ComplexTransform::kNone;
// Legacy: Assume MultiplyAdd only since we do not use this tag type in 3.0
using MathOperator = cutlass::arch::OpMultiplyAdd;
using OperatorClass = cutlass::detail::get_operator_class_t<typename CollectiveMainloop::TiledMma>;
using ArchTag = typename GemmKernel::ArchTag;
// NOTE: Assume identity swizzle for now
using ThreadblockSwizzle = cutlass::gemm::threadblock::GemmIdentityThreadblockSwizzle<>;
// Assume TiledMma's ShapeMNK is the same as 2.x's ThreadblockShape
using ThreadblockShape = cutlass::gemm::GemmShape<
cute::size<0>(TileShape{}),
cute::size<1>(TileShape{}),
cute::size<2>(TileShape{})>;
using ClusterShape = cutlass::gemm::GemmShape<
cute::size<0>(typename GemmKernel::DispatchPolicy::ClusterShape{}),
cute::size<1>(typename GemmKernel::DispatchPolicy::ClusterShape{}),
cute::size<2>(typename GemmKernel::DispatchPolicy::ClusterShape{})>;
// Instruction shape is easy too, since we get that directly from our TiledMma's atom shape
using InstructionShape = cutlass::gemm::GemmShape<
cute::size<0>(typename CollectiveMainloop::TiledMma::AtomShape_MNK{}),
cute::size<1>(typename CollectiveMainloop::TiledMma::AtomShape_MNK{}),
cute::size<2>(typename CollectiveMainloop::TiledMma::AtomShape_MNK{})>;
// Legacy: provide a correct warp count, but no reliable warp shape
static int const kThreadCount = GemmKernel::MaxThreadsPerBlock;
// Warp shape is not a primary API type in 3.x
// But we can best approximate it by inspecting the TiledMma
// For this, we make the assumption that we always have 4 warps along M, and rest along N, none along K
// We also always round up the warp count to 4 if the tiled mma is smaller than 128 threads
static constexpr int WarpsInMma = cute::max(4, CUTE_STATIC_V(cute::size(typename GemmKernel::TiledMma{})) / 32);
static constexpr int WarpsInMmaM = 4;
static constexpr int WarpsInMmaN = cute::ceil_div(WarpsInMma, WarpsInMmaM);
using WarpCount = cutlass::gemm::GemmShape<WarpsInMmaM, WarpsInMmaN, 1>;
using WarpShape = cutlass::gemm::GemmShape<
CUTE_STATIC_V(cute::tile_size<0>(typename CollectiveMainloop::TiledMma{})) / WarpsInMmaM,
CUTE_STATIC_V(cute::tile_size<1>(typename CollectiveMainloop::TiledMma{})) / WarpsInMmaN,
CUTE_STATIC_V(cute::tile_size<2>(typename CollectiveMainloop::TiledMma{}))>;
static int constexpr kStages = detail::stages_member(typename CollectiveMainloop::DispatchPolicy{});
// Inspect TiledCopy for A and B to compute the alignment size
static int constexpr kAlignmentA = cutlass::detail::get_alignment_count_from_gmem_tiled_copy<
typename CollectiveMainloop::GmemTiledCopyA, ElementA, typename CollectiveMainloop::TiledMma::ValTypeA>();
static int constexpr kAlignmentB = cutlass::detail::get_alignment_count_from_gmem_tiled_copy<
typename CollectiveMainloop::GmemTiledCopyB, ElementB, typename CollectiveMainloop::TiledMma::ValTypeB>();
static int constexpr kAlignmentC = cutlass::detail::get_alignment_count_from_gmem_tiled_copy<
typename CollectiveEpilogue::GmemTiledCopyC, ElementC>();
static int constexpr kAlignmentD = cutlass::detail::get_alignment_count_from_gmem_tiled_copy<
typename CollectiveEpilogue::GmemTiledCopyD, ElementD>();
using EpilogueOutputOp = typename CollectiveEpilogue::ThreadEpilogueOp;
// Split-K preserves splits that are 128b aligned
static int constexpr kSplitKAlignment = cute::max(
128 / cutlass::sizeof_bits<ElementA>::value, 128 / cutlass::sizeof_bits<ElementB>::value);
/// Argument structure: User API
using Arguments = typename GemmKernel::Arguments;
/// Argument structure: Kernel API
using Params = typename GemmKernel::Params;
private:
/// Kernel API parameters object
Params params_;
public:
/// Access the Params structure
Params const& params() const {
return params_;
}
/// Determines whether the GEMM can execute the given problem.
static cutlass::Status
can_implement(Arguments const& args) {
if (GemmKernel::can_implement(args)) {
return cutlass::Status::kSuccess;
}
else {
return cutlass::Status::kInvalid;
}
}
/// Gets the workspace size
static size_t
get_workspace_size(Arguments const& args) {
size_t workspace_bytes = 0;
if (args.mode == cutlass::gemm::GemmUniversalMode::kGemmSplitKParallel) {
workspace_bytes += sizeof(int) * size_t(cute::size<0>(TileShape{})) * size_t(cute::size<1>(TileShape{}));
}
workspace_bytes += GemmKernel::get_workspace_size(args);
CUTLASS_TRACE_HOST(" workspace_bytes: " << workspace_bytes);
return workspace_bytes;
}
/// Computes the grid shape
static dim3
get_grid_shape(Arguments const& args, void* workspace = nullptr) {
auto tmp_params = GemmKernel::to_underlying_arguments(args, workspace);
return GemmKernel::get_grid_shape(tmp_params);
}
/// Computes the grid shape
static dim3
get_grid_shape(Params const& params) {
return GemmKernel::get_grid_shape(params);
}
/// Computes the maximum number of active blocks per multiprocessor
static int maximum_active_blocks(int /* smem_capacity */ = -1) {
CUTLASS_TRACE_HOST("GemmUniversal::maximum_active_blocks()");
int max_active_blocks = -1;
int smem_size = GemmKernel::SharedStorageSize;
// first, account for dynamic smem capacity if needed
cudaError_t result;
if (smem_size >= (48 << 10)) {
CUTLASS_TRACE_HOST(" Setting smem size to " << smem_size);
result = cudaFuncSetAttribute(
cutlass::device_kernel<GemmKernel>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
smem_size);
if (cudaSuccess != result) {
result = cudaGetLastError(); // to clear the error bit
CUTLASS_TRACE_HOST(
" cudaFuncSetAttribute() returned error: "
<< cudaGetErrorString(result));
return -1;
}
}
// query occupancy after setting smem size
result = cudaOccupancyMaxActiveBlocksPerMultiprocessor(
&max_active_blocks,
cutlass::device_kernel<GemmKernel>,
GemmKernel::MaxThreadsPerBlock,
smem_size);
if (cudaSuccess != result) {
result = cudaGetLastError(); // to clear the error bit
CUTLASS_TRACE_HOST(
" cudaOccupancyMaxActiveBlocksPerMultiprocessor() returned error: "
<< cudaGetErrorString(result));
return -1;
}
CUTLASS_TRACE_HOST(" max_active_blocks: " << max_active_blocks);
return max_active_blocks;
}
/// Initializes GEMM state from arguments.
cutlass::Status
initialize(
Arguments const& args,
void* workspace = nullptr,
cutlass::CudaHostAdapter* cuda_adapter = nullptr) {
// Initialize the workspace
cutlass::Status status = GemmKernel::initialize_workspace(args, workspace, cuda_adapter);
if (status != cutlass::Status::kSuccess) {
return status;
}
// Initialize the Params structure
params_ = GemmKernel::to_underlying_arguments(args, workspace);
// Don't set the function attributes - require the CudaHostAdapter to set it.
if constexpr (kEnableCudaHostAdapter) {
CUTLASS_ASSERT(cuda_adapter);
return cutlass::Status::kSuccess;
}
else {
//
// Account for dynamic smem capacity if needed
//
int smem_size = GemmKernel::SharedStorageSize;
CUTLASS_ASSERT(cuda_adapter == nullptr);
if (smem_size >= (48 << 10)) {
CUTLASS_TRACE_HOST(" Setting smem size to " << smem_size);
cudaError_t result = cudaFuncSetAttribute(
cutlass::device_kernel<GemmKernel>,
cudaFuncAttributeMaxDynamicSharedMemorySize,
smem_size);
if (cudaSuccess != result) {
result = cudaGetLastError(); // to clear the error bit
CUTLASS_TRACE_HOST(" cudaFuncSetAttribute() returned error: " << cudaGetErrorString(result));
return cutlass::Status::kErrorInternal;
}
}
}
return cutlass::Status::kSuccess;
}
/// Update API is preserved in 3.0, but does not guarantee a lightweight update of params.
cutlass::Status
update(Arguments const& args, void* workspace = nullptr) {
CUTLASS_TRACE_HOST("GemmUniversal()::update() - workspace: " << workspace);
size_t workspace_bytes = get_workspace_size(args);
if (workspace_bytes > 0 && nullptr == workspace) {
return cutlass::Status::kErrorWorkspaceNull;
}
params_ = GemmKernel::to_underlying_arguments(args, workspace);
return cutlass::Status::kSuccess;
}
/// Primary run() entry point API that is static allowing users to create and manage their own params.
/// Supplied params struct must be construct by calling GemmKernel::to_underlying_arguments()
static cutlass::Status
run(Params& params,
cutlass::CudaHostAdapter *cuda_adapter = nullptr,
bool launch_with_pdl = false) {
CUTLASS_TRACE_HOST("GemmUniversal::run()");
dim3 const block = GemmKernel::get_block_shape();
dim3 const grid = get_grid_shape(params);
// configure smem size and carveout
int smem_size = GemmKernel::SharedStorageSize;
cutlass::Status launch_result{ cutlass::Status::kSuccess };
// Use extended launch API only for mainloops that use it
if constexpr (GemmKernel::ArchTag::kMinComputeCapability >= 90) {
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
CUTLASS_TRACE_HOST("GemmUniversal::run: Use extended launch API");
#endif
[[maybe_unused]] constexpr bool is_static_1x1x1 =
cute::is_static_v<typename GemmKernel::DispatchPolicy::ClusterShape> and
cute::size(typename GemmKernel::DispatchPolicy::ClusterShape{}) == 1;
[[maybe_unused]] dim3 cluster(cute::size<0>(typename GemmKernel::DispatchPolicy::ClusterShape{}),
cute::size<1>(typename GemmKernel::DispatchPolicy::ClusterShape{}),
cute::size<2>(typename GemmKernel::DispatchPolicy::ClusterShape{}));
// Dynamic cluster support
[[maybe_unused]] dim3 fallback_cluster = dim3{0,0,0};
if constexpr (GemmKernel::ArchTag::kMinComputeCapability == 100
|| GemmKernel::ArchTag::kMinComputeCapability == 101
|| GemmKernel::ArchTag::kMinComputeCapability == 103
) {
if constexpr (!cute::is_static_v<typename GemmKernel::DispatchPolicy::ClusterShape>) {
if constexpr (detail::IsDistGemmKernel<GemmKernel>::value) {
fallback_cluster = params.base.hw_info.cluster_shape_fallback;
cluster = params.base.hw_info.cluster_shape;
} else {
fallback_cluster = params.hw_info.cluster_shape_fallback;
cluster = params.hw_info.cluster_shape;
}
}
}
[[maybe_unused]] void* kernel_params[] = {¶ms};
if constexpr (kEnableCudaHostAdapter) {
//
// Use the cuda host adapter
//
CUTLASS_ASSERT(cuda_adapter);
if (cuda_adapter) {
if (launch_with_pdl) {
CUTLASS_TRACE_HOST(
"GemmUniversal::run() does not support launching with PDL and a custom cuda adapter.");
return cutlass::Status::kErrorInternal;
}
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
CUTLASS_TRACE_HOST("GemmUniversal::run: Launching kernel with CUDA host adapter");
#endif
if constexpr (is_static_1x1x1) {
launch_result = cuda_adapter->launch(grid,
block,
smem_size,
nullptr,
kernel_params,
0);
}
else {
launch_result = cuda_adapter->launch(grid,
cluster,
fallback_cluster,
block,
smem_size,
nullptr,
kernel_params,
0);
}
}
else {
CUTLASS_TRACE_HOST("GemmUniversal::run: kEnableCudaHostAdapter is true, but CUDA host adapter is null");
return cutlass::Status::kErrorInternal;
}
}
else {
CUTLASS_ASSERT(cuda_adapter == nullptr);
[[maybe_unused]] void const* kernel = (void const*) cutlass::device_kernel<GemmKernel>;
static constexpr bool kClusterLaunch = GemmKernel::ArchTag::kMinComputeCapability == 90;
if constexpr (kClusterLaunch) {
if constexpr (is_static_1x1x1) {
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
CUTLASS_TRACE_HOST("GemmUniversal::run: Launching static 1x1x1 kernel");
#endif
launch_result = cutlass::kernel_launch<GemmKernel>(
grid, block, smem_size, nullptr, params, launch_with_pdl);
if (launch_result != cutlass::Status::kSuccess) {
CUTLASS_TRACE_HOST("GemmUniversal::run: cutlass::kernel_launch reports failure");
}
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
else {
CUTLASS_TRACE_HOST("GemmUniversal::run: cutlass::kernel_launch reports success");
}
#endif
}
else {
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
CUTLASS_TRACE_HOST("GemmUniversal::run: Launching dynamic cluster kernel");
#endif
launch_result = cutlass::ClusterLauncher::launch(
grid, cluster, block, smem_size, nullptr, kernel, kernel_params, launch_with_pdl);
}
}
else {
if constexpr (GemmKernel::ArchTag::kMinComputeCapability == 100
|| GemmKernel::ArchTag::kMinComputeCapability == 101
|| GemmKernel::ArchTag::kMinComputeCapability == 120
|| GemmKernel::ArchTag::kMinComputeCapability == 103
) {
if constexpr (is_static_1x1x1) {
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
CUTLASS_TRACE_HOST("GemmUniversal::run: Launching static 1x1x1 kernel");
#endif
launch_result = cutlass::kernel_launch<GemmKernel>(grid, block, smem_size, nullptr, params, launch_with_pdl);
if (launch_result != cutlass::Status::kSuccess) {
CUTLASS_TRACE_HOST("GemmUniversal::run: cutlass::kernel_launch reports failure");
}
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
else {
CUTLASS_TRACE_HOST("GemmUniversal::run: cutlass::kernel_launch reports success");
}
#endif
}
else {
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
CUTLASS_TRACE_HOST("GemmUniversal::run: Launching kernel with fall-back cluster");
#endif
launch_result = cutlass::ClusterLauncher::launch_with_fallback_cluster(
grid,
cluster,
fallback_cluster,
block,
smem_size,
nullptr,
kernel,
kernel_params,
launch_with_pdl);
}
}
}
}
}
else {
launch_result = cutlass::Status::kSuccess;
cutlass::arch::synclog_setup();
if constexpr (kEnableCudaHostAdapter) {
CUTLASS_ASSERT(cuda_adapter);
if (cuda_adapter) {
void* kernel_params[] = {¶ms};
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
CUTLASS_TRACE_HOST("GemmUniversal::run: Launching kernel with CUDA host adapter");
#endif
launch_result = cuda_adapter->launch(
grid, block, smem_size, nullptr, kernel_params, 0
);
}
else {
CUTLASS_TRACE_HOST("GemmUniversal::run: CUDA host adapter is null");
return cutlass::Status::kErrorInternal;
}
}
else {
CUTLASS_ASSERT(cuda_adapter == nullptr);
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
CUTLASS_TRACE_HOST("GemmUniversal::run: Launching kernel with cutlass::kernel_launch");
#endif
launch_result = cutlass::kernel_launch<GemmKernel>(
grid, block, smem_size, nullptr, params, launch_with_pdl);
if (launch_result != cutlass::Status::kSuccess) {
CUTLASS_TRACE_HOST("GemmUniversal::run: cutlass::kernel_launch reports failure");
}
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
else {
CUTLASS_TRACE_HOST("GemmUniversal::run: cutlass::kernel_launch reports success");
}
#endif
}
}
cudaError_t result = cudaGetLastError();
if (cudaSuccess == result && cutlass::Status::kSuccess == launch_result) {
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
CUTLASS_TRACE_HOST("GemmUniversal::run: cudaGetLastError reports success");
#endif
return cutlass::Status::kSuccess;
}
else {
CUTLASS_TRACE_HOST(" Kernel launch failed. Reason: " << result);
return cutlass::Status::kErrorInternal;
}
}
//
// Non-static launch overloads that first create and set the internal params struct of this kernel handle.
//
/// Launches the kernel after first constructing Params internal state from supplied arguments.
cutlass::Status
run(
Arguments const& args,
void* workspace = nullptr,
cutlass::CudaHostAdapter *cuda_adapter = nullptr,
bool launch_with_pdl = false
) {
cutlass::Status status = initialize(args, workspace, nullptr, cuda_adapter);
if (cutlass::Status::kSuccess == status) {
status = run(params_, cuda_adapter, launch_with_pdl);
}
return status;
}
/// Launches the kernel after first constructing Params internal state from supplied arguments.
cutlass::Status
operator()(
Arguments const& args,
void* workspace = nullptr,
cutlass::CudaHostAdapter *cuda_adapter = nullptr,
bool launch_with_pdl = false) {
return run(args, workspace, cuda_adapter, launch_with_pdl);
}
/// Overload that allows a user to re-launch the same kernel without updating internal params struct.
cutlass::Status
run(
cutlass::CudaHostAdapter *cuda_adapter = nullptr,
bool launch_with_pdl = false) {
return run(params_, cuda_adapter, launch_with_pdl);
}
/// Overload that allows a user to re-launch the same kernel without updating internal params struct.
cutlass::Status
operator()(cutlass::CudaHostAdapter *cuda_adapter = nullptr, bool launch_with_pdl = false) {
return run(params_, cuda_adapter, launch_with_pdl);
}
};
////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::gemm::device
////////////////////////////////////////////////////////////////////////////////
// ========== sm100_gemm_tma_warpspecialized.hpp ==========
#include <cutlass/cutlass.h>
#include <cutlass/workspace.h>
#include <cutlass/kernel_hardware_info.hpp>
#include <cutlass/detail/cluster.hpp>
#include <cutlass/arch/grid_dependency_control.h>
#include <cutlass/fast_math.h>
#include <cute/arch/cluster_sm90.hpp>
#include <cutlass/arch/arch.h>
#include <cutlass/arch/barrier.h>
#include <cutlass/arch/reg_reconfig.h>
#include <cutlass/gemm/gemm.h>
#include <cutlass/gemm/dispatch_policy.hpp>
#include <cutlass/detail/mainloop_fusion_helper_scale_factor.hpp>
#include <cutlass/gemm/kernel/sm100_tile_scheduler.hpp>
#include <cutlass/pipeline/pipeline.hpp>
#include <cutlass/detail/sm100_tmem_helper.hpp>
#include <cute/tensor.hpp>
#include <cute/arch/tmem_allocator_sm100.hpp>
#include <cute/atom/mma_atom.hpp>
///////////////////////////////////////////////////////////////////////////////
namespace gpu_mode {
using namespace cute;
///////////////////////////////////////////////////////////////////////////////
template <
class ProblemShape_,
class CollectiveMainloop_,
class CollectiveEpilogue_,
class TileSchedulerTag_
>
class GemmUniversalCustom : cutlass::gemm::kernel::GemmUniversal<
ProblemShape_,
CollectiveMainloop_,
CollectiveEpilogue_,
TileSchedulerTag_
>
{
public:
//
// Type Aliases
//
using ProblemShape = ProblemShape_;
static_assert(rank(ProblemShape{}) == 3 or rank(ProblemShape{}) == 4,
"ProblemShape{} should be <M,N,K> or <M,N,K,L>");
// Mainloop derived types
using CollectiveMainloop = CollectiveMainloop_;
using TileShape = typename CollectiveMainloop::TileShape;
using TiledMma = typename CollectiveMainloop::TiledMma;
using ArchTag = typename CollectiveMainloop::ArchTag;
using ElementA = typename CollectiveMainloop::ElementA;
using StrideA = typename CollectiveMainloop::StrideA;
using ElementB = typename CollectiveMainloop::ElementB;
using StrideB = typename CollectiveMainloop::StrideB;
using LayoutSFA = typename cutlass::detail::LayoutSFAType<CollectiveMainloop>::type;
using LayoutSFB = typename cutlass::detail::LayoutSFBType<CollectiveMainloop>::type;
using ElementSF = typename cutlass::detail::ElementSFType<CollectiveMainloop>::type;
using DispatchPolicy = typename CollectiveMainloop::DispatchPolicy;
using ElementAccumulator = typename CollectiveMainloop::ElementAccumulator;
using ClusterShape = typename DispatchPolicy::ClusterShape;
using MainloopArguments = typename CollectiveMainloop::Arguments;
using MainloopParams = typename CollectiveMainloop::Params;
static_assert(ArchTag::kMinComputeCapability >= 100);
// Epilogue derived types
using CollectiveEpilogue = CollectiveEpilogue_;
using EpilogueTile = typename CollectiveEpilogue::EpilogueTile;
using ElementC = typename CollectiveEpilogue::ElementC;
using StrideC = typename CollectiveEpilogue::StrideC;
using ElementD = typename CollectiveEpilogue::ElementD;
using StrideD = typename CollectiveEpilogue::StrideD;
using EpilogueArguments = typename CollectiveEpilogue::Arguments;
using EpilogueParams = typename CollectiveEpilogue::Params;
static constexpr bool IsComplex = CollectiveEpilogue::NumAccumulatorMtxs == 2;
// CLC pipeline depth
// determines how many waves (stages-1) a warp can race ahead
static constexpr uint32_t SchedulerPipelineStageCount = DispatchPolicy::Schedule::SchedulerPipelineStageCount;
static constexpr uint32_t AccumulatorPipelineStageCount = DispatchPolicy::Schedule::AccumulatorPipelineStageCount;
static constexpr bool IsOverlappingAccum = DispatchPolicy::IsOverlappingAccum;
// TileID scheduler
// Get Blk and Scheduling tile shapes
using AtomThrShapeMNK = typename CollectiveMainloop::AtomThrShapeMNK;
using CtaShape_MNK = typename CollectiveMainloop::CtaShape_MNK;
using TileSchedulerTag = TileSchedulerTag_;
using TileScheduler = typename cutlass::gemm::kernel::detail::TileSchedulerSelector<
TileSchedulerTag, ArchTag, CtaShape_MNK, ClusterShape, SchedulerPipelineStageCount>::Scheduler;
using TileSchedulerArguments = typename TileScheduler::Arguments;
using TileSchedulerParams = typename TileScheduler::Params;
static constexpr bool IsSchedDynamicPersistent = TileScheduler::IsDynamicPersistent;
static constexpr bool IsDynamicCluster = not cute::is_static_v<ClusterShape>;
static constexpr bool IsGdcEnabled = cutlass::arch::IsGdcGloballyEnabled;
// Warp specialization thread count per threadblock
static constexpr uint32_t NumSchedThreads = cutlass::NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumMMAThreads = cutlass::NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumMainloopLoadThreads = cutlass::NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumEpilogueLoadThreads = cutlass::NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumEpilogueThreads = CollectiveEpilogue::ThreadCount;
static constexpr uint32_t NumEpilogueWarps = NumEpilogueThreads / cutlass::NumThreadsPerWarp;
static constexpr uint32_t MaxThreadsPerBlock = NumSchedThreads +
NumMainloopLoadThreads + NumMMAThreads +
NumEpilogueLoadThreads + NumEpilogueThreads;
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
static constexpr uint32_t NumEpilogueSubTiles = CollectiveEpilogue::get_load_pipe_increment(CtaShape_MNK{});
// Fixup performed for split-/strxxm-K is done across warps in different CTAs
// at epilogue subtile granularity. Thus, there must be one barrier per sub-tile per
// epilogue warp.
static constexpr uint32_t NumFixupBarriers = 1;
static constexpr uint32_t CLCResponseSize = sizeof(typename TileScheduler::CLCResponse);
// Pipeline and pipeline state types
using MainloopPipeline = typename CollectiveMainloop::MainloopPipeline;
using MainloopPipelineState = typename CollectiveMainloop::MainloopPipelineState;
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
using EpiLoadPipelineState = typename CollectiveEpilogue::LoadPipelineState;
using EpiStorePipeline = typename CollectiveEpilogue::StorePipeline;
using EpiStorePipelineState = typename CollectiveEpilogue::StorePipelineState;
using LoadOrderBarrier = cutlass::OrderedSequenceBarrier<1,2>;
using AccumulatorPipeline = cutlass::PipelineUmmaAsync<AccumulatorPipelineStageCount, AtomThrShapeMNK>;
using AccumulatorPipelineState = typename AccumulatorPipeline::PipelineState;
using CLCPipeline = cutlass::PipelineCLCFetchAsync<SchedulerPipelineStageCount, ClusterShape>;
using CLCPipelineState = typename CLCPipeline::PipelineState;
using CLCThrottlePipeline = cutlass::PipelineAsync<SchedulerPipelineStageCount>;
using CLCThrottlePipelineState = typename CLCThrottlePipeline::PipelineState;
using TmemAllocator = cute::conditional_t<cute::size(cute::shape<0>(typename TiledMma::ThrLayoutVMNK{})) == 1,
cute::TMEM::Allocator1Sm, cute::TMEM::Allocator2Sm>;
// Kernel level shared memory storage
struct SharedStorage {
struct PipelineStorage : cute::aligned_struct<16, _1> {
using MainloopPipelineStorage = typename CollectiveMainloop::PipelineStorage;
using EpiLoadPipelineStorage = typename CollectiveEpilogue::PipelineStorage;
using LoadOrderBarrierStorage = typename LoadOrderBarrier::SharedStorage;
using CLCPipelineStorage = typename CLCPipeline::SharedStorage;
using AccumulatorPipelineStorage = typename AccumulatorPipeline::SharedStorage;
using CLCThrottlePipelineStorage = typename CLCThrottlePipeline::SharedStorage;
alignas(16) MainloopPipelineStorage mainloop;
alignas(16) EpiLoadPipelineStorage epi_load;
alignas(16) LoadOrderBarrierStorage load_order;
alignas(16) CLCPipelineStorage clc;
alignas(16) AccumulatorPipelineStorage accumulator;
alignas(16) CLCThrottlePipelineStorage clc_throttle;
alignas(16) cutlass::arch::ClusterBarrier tmem_dealloc;
} pipelines;
alignas(16) typename TileScheduler::CLCResponse clc_response[SchedulerPipelineStageCount];
uint32_t tmem_base_ptr;
struct TensorStorage : cute::aligned_struct<128, _1> {
using EpilogueTensorStorage = typename CollectiveEpilogue::TensorStorage;
using MainloopTensorStorage = typename CollectiveMainloop::TensorStorage;
EpilogueTensorStorage epilogue;
MainloopTensorStorage mainloop;
} tensors;
};
static constexpr int SharedStorageSize = sizeof(SharedStorage);
// Host facing host arguments
struct Arguments {
cutlass::gemm::GemmUniversalMode mode{};
ProblemShape problem_shape{};
MainloopArguments mainloop{};
EpilogueArguments epilogue{};
cutlass::KernelHardwareInfo hw_info{};
TileSchedulerArguments scheduler{};
};
// Kernel device entry point API
struct Params {
cutlass::gemm::GemmUniversalMode mode{};
ProblemShape problem_shape{};
MainloopParams mainloop{};
EpilogueParams epilogue{};
TileSchedulerParams scheduler{};
cutlass::KernelHardwareInfo hw_info{};
};
enum class WarpCategory : int32_t {
MMA = 0,
Sched = 1,
MainloopLoad = 2,
EpilogueLoad = 3,
Epilogue = 4
};
struct IsParticipant {
uint32_t mma = false;
uint32_t sched = false;
uint32_t main_load = false;
uint32_t epi_load = false;
uint32_t epilogue = false;
};
//
// Methods
//
// Convert to underlying arguments.
static
Params
to_underlying_arguments(Arguments const& args, void* workspace) {
(void) workspace;
auto problem_shape = args.problem_shape;
auto problem_shape_MNKL = append<4>(problem_shape, 1);
// Get SM count if needed, otherwise use user supplied SM count
int sm_count = args.hw_info.sm_count;
if (sm_count != 0) {
CUTLASS_TRACE_HOST(" WARNING: SM100 tile scheduler does not allow for user specified SM counts.\\n"
" To restrict a kernel's resource usage, consider using CUDA driver APIs instead (green contexts).");
}
CUTLASS_TRACE_HOST("to_underlying_arguments(): Setting persistent grid SM count to " << sm_count);
// Calculate workspace pointers
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
size_t workspace_offset = 0;
// Epilogue
void* epilogue_workspace = workspace_ptr + workspace_offset;
workspace_offset += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue);
workspace_offset = cutlass::round_nearest(workspace_offset, cutlass::MinWorkspaceAlignment);
void* mainloop_workspace = nullptr;
// Tile scheduler
void* scheduler_workspace = workspace_ptr + workspace_offset;
workspace_offset += TileScheduler::template get_workspace_size<ProblemShape, ElementAccumulator>(
args.scheduler, args.problem_shape, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs);
workspace_offset = cutlass::round_nearest(workspace_offset, cutlass::MinWorkspaceAlignment);
return {
args.mode,
args.problem_shape,
CollectiveMainloop::to_underlying_arguments(args.problem_shape, args.mainloop, mainloop_workspace, args.hw_info),
CollectiveEpilogue::to_underlying_arguments(args.problem_shape, args.epilogue, epilogue_workspace),
TileScheduler::to_underlying_arguments(
problem_shape_MNKL, TileShape{}, AtomThrShapeMNK{}, ClusterShape{},
args.hw_info, args.scheduler, scheduler_workspace
)
,args.hw_info
};
}
static bool
can_implement(Arguments const& args) {
bool implementable = (args.mode == cutlass::gemm::GemmUniversalMode::kGemm) or
(args.mode == cutlass::gemm::GemmUniversalMode::kBatched && rank(ProblemShape{}) == 4);
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Arguments or Problem Shape don't meet the requirements.\\n");
return implementable;
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
if constexpr (IsDynamicCluster) {
static constexpr int MaxClusterSize = 16;
implementable &= size(args.hw_info.cluster_shape) <= MaxClusterSize;
implementable &= size(args.hw_info.cluster_shape_fallback) <= MaxClusterSize;
implementable &= cutlass::detail::preferred_cluster_can_implement<AtomThrShapeMNK>(args.hw_info.cluster_shape, args.hw_info.cluster_shape_fallback);
}
constexpr bool IsBlockscaled = !cute::is_void_v<ElementSF>;
if constexpr (IsBlockscaled) {
if constexpr (IsDynamicCluster) {
implementable &= cutlass::detail::preferred_cluster_can_implement<AtomThrShapeMNK>(args.hw_info.cluster_shape, args.hw_info.cluster_shape_fallback);
// Special cluster shape check for scale factor multicasts. Due to limited size of scale factors, we can't multicast among
// more than 4 CTAs
implementable &= (args.hw_info.cluster_shape.x <= 4 && args.hw_info.cluster_shape.y <= 4 &&
args.hw_info.cluster_shape_fallback.x <= 4 && args.hw_info.cluster_shape_fallback.y <= 4);
}
else {
// Special cluster shape check for scale factor multicasts. Due to limited size of scale factors, we can't multicast among
// more than 4 CTAs
implementable &= ((size<0>(ClusterShape{}) <= 4) && (size<1>(ClusterShape{}) <= 4));
}
}
return implementable;
}
static size_t
get_workspace_size(Arguments const& args) {
size_t workspace_size = 0;
// Epilogue
workspace_size += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue);
workspace_size = cutlass::round_nearest(workspace_size, cutlass::MinWorkspaceAlignment);
// Tile scheduler
workspace_size += TileScheduler::template get_workspace_size<ProblemShape, ElementAccumulator>(
args.scheduler, args.problem_shape, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs);
workspace_size = cutlass::round_nearest(workspace_size, cutlass::MinWorkspaceAlignment);
return workspace_size;
}
static cutlass::Status
initialize_workspace(Arguments const& args, void* workspace = nullptr,
cutlass::CudaHostAdapter* cuda_adapter = nullptr) {
cutlass::Status status = cutlass::Status::kSuccess;
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
size_t workspace_offset = 0;
// Epilogue
status = CollectiveEpilogue::initialize_workspace(args.problem_shape, args.epilogue, workspace_ptr + workspace_offset, nullptr, cuda_adapter);
workspace_offset += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue);
workspace_offset = cutlass::round_nearest(workspace_offset, cutlass::MinWorkspaceAlignment);
if (status != cutlass::Status::kSuccess) {
return status;
}
// Tile scheduler
status = TileScheduler::template initialize_workspace<ProblemShape, ElementAccumulator>(
args.scheduler, workspace_ptr + workspace_offset, nullptr, args.problem_shape, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs, cuda_adapter);
workspace_offset += TileScheduler::template get_workspace_size<ProblemShape, ElementAccumulator>(
args.scheduler, args.problem_shape, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs);
workspace_offset = cutlass::round_nearest(workspace_offset, cutlass::MinWorkspaceAlignment);
if (status != cutlass::Status::kSuccess) {
return status;
}
return status;
}
// Computes the kernel launch grid shape based on runtime parameters
static dim3
get_grid_shape(Params const& params) {
// NOTE cluster_shape here is the major cluster shape, not fallback one
auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, params.hw_info.cluster_shape);
auto problem_shape_MNKL = append<4>(params.problem_shape, _1{});
return TileScheduler::get_grid_shape(
params.scheduler,
problem_shape_MNKL,
TileShape{},
AtomThrShapeMNK{},
cluster_shape,
params.hw_info);
}
static dim3
get_block_shape() {
return dim3(MaxThreadsPerBlock, 1, 1);
}
CUTLASS_DEVICE
void
operator() (Params const& params, char* smem_buf) {
using namespace cute;
using X = Underscore;
static_assert(SharedStorageSize <= cutlass::arch::sm100_smem_capacity_bytes, "SMEM usage exceeded capacity.");
// Separate out problem shape for convenience
// Optionally append 1s until problem shape is rank-4 in case its is only rank-3 (MNK)
auto problem_shape_MNKL = append<4>(params.problem_shape, _1{});
auto [M,N,K,L] = problem_shape_MNKL;
// Account for more than one epilogue warp
int warp_idx = cutlass::canonical_warp_idx_sync();
WarpCategory warp_category = warp_idx < static_cast<int>(WarpCategory::Epilogue) ? WarpCategory(warp_idx)
: WarpCategory::Epilogue;
uint32_t lane_predicate = cute::elect_one_sync();
auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{});
int cluster_size = size(cluster_shape);
uint32_t cta_rank_in_cluster = cute::block_rank_in_cluster();
bool is_first_cta_in_cluster = cta_rank_in_cluster == 0;
int cta_coord_v = cta_rank_in_cluster % size<0>(typename TiledMma::AtomThrID{});
bool is_mma_leader_cta = cta_coord_v == 0;
constexpr bool has_mma_peer_cta = size(AtomThrShapeMNK{}) == 2;
[[maybe_unused]] uint32_t mma_peer_cta_rank = has_mma_peer_cta ? cta_rank_in_cluster ^ 1 : cta_rank_in_cluster;
// Kernel level shared memory storage
SharedStorage& shared_storage = *reinterpret_cast<SharedStorage*>(smem_buf);
// In a warp specialized kernel, collectives expose data movement and compute operations separately
CollectiveMainloop collective_mainloop(params.mainloop, cluster_shape, cta_rank_in_cluster);
CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.epilogue);
// Issue Tma Descriptor Prefetch from a single thread
if ((warp_category == WarpCategory::Sched) && lane_predicate) {
collective_mainloop.prefetch_tma_descriptors();
}
if ((warp_category == WarpCategory::EpilogueLoad) && lane_predicate) {
collective_epilogue.prefetch_tma_descriptors(params.epilogue);
}
// Do we load source tensor C or other aux inputs
bool is_epi_load_needed = collective_epilogue.is_producer_load_needed();
IsParticipant is_participant = {
(warp_category == WarpCategory::MMA), // mma
(warp_category == WarpCategory::Sched) && is_first_cta_in_cluster, // sched
(warp_category == WarpCategory::MainloopLoad), // main_load
(warp_category == WarpCategory::EpilogueLoad) && is_epi_load_needed, // epi_load
(warp_category == WarpCategory::Epilogue) // epilogue
};
// Mainloop Load pipeline
typename MainloopPipeline::Params mainloop_pipeline_params;
if (WarpCategory::MainloopLoad == warp_category) {
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Producer;
}
if (WarpCategory::MMA == warp_category) {
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Consumer;
}
mainloop_pipeline_params.is_leader = lane_predicate && is_mma_leader_cta && is_participant.main_load;
mainloop_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
mainloop_pipeline_params.initializing_warp = 0;
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop,
mainloop_pipeline_params,
cluster_shape,
cute::true_type{}, // Perform barrier init
cute::false_type{}); // Delay mask calculation
// Epilogue Load pipeline
typename EpiLoadPipeline::Params epi_load_pipeline_params;
if (WarpCategory::EpilogueLoad == warp_category) {
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Producer;
}
if (WarpCategory::Epilogue == warp_category) {
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Consumer;
}
epi_load_pipeline_params.dst_blockid = cta_rank_in_cluster;
epi_load_pipeline_params.producer_arv_count = NumEpilogueLoadThreads;
epi_load_pipeline_params.consumer_arv_count = NumEpilogueThreads;
epi_load_pipeline_params.transaction_bytes = CollectiveEpilogue::TmaTransactionBytes;
epi_load_pipeline_params.initializing_warp = 1;
EpiLoadPipeline epi_load_pipeline(shared_storage.pipelines.epi_load, epi_load_pipeline_params);
// Epilogue Store pipeline
typename EpiStorePipeline::Params epi_store_pipeline_params;
epi_store_pipeline_params.always_wait = true;
EpiStorePipeline epi_store_pipeline(epi_store_pipeline_params);
// Load order barrier
typename LoadOrderBarrier::Params load_order_barrier_params;
load_order_barrier_params.group_id = (warp_category == WarpCategory::MainloopLoad) ? 0 : 1;
load_order_barrier_params.group_size = NumMainloopLoadThreads;
load_order_barrier_params.initializing_warp = 3;
LoadOrderBarrier load_order_barrier(shared_storage.pipelines.load_order, load_order_barrier_params);
// CLC pipeline
typename CLCPipeline::Params clc_pipeline_params;
if (WarpCategory::Sched == warp_category) {
clc_pipeline_params.role = CLCPipeline::ThreadCategory::ProducerConsumer;
}
else {
clc_pipeline_params.role = CLCPipeline::ThreadCategory::Consumer;
}
clc_pipeline_params.producer_blockid = 0;
clc_pipeline_params.producer_arv_count = 1;
clc_pipeline_params.consumer_arv_count = NumSchedThreads + cluster_size *
(NumMainloopLoadThreads + NumEpilogueThreads + NumMMAThreads);
if (is_epi_load_needed) {
clc_pipeline_params.consumer_arv_count += cluster_size * NumEpilogueLoadThreads;
}
clc_pipeline_params.transaction_bytes = CLCResponseSize;
clc_pipeline_params.initializing_warp = 4;
CLCPipeline clc_pipeline(shared_storage.pipelines.clc, clc_pipeline_params, cluster_shape);
// Mainloop-Epilogue pipeline
typename AccumulatorPipeline::Params accumulator_pipeline_params;
if (WarpCategory::MMA == warp_category) {
accumulator_pipeline_params.role = AccumulatorPipeline::ThreadCategory::Producer;
}
if (WarpCategory::Epilogue == warp_category) {
accumulator_pipeline_params.role = AccumulatorPipeline::ThreadCategory::Consumer;
}
// Only one producer thread arrives on this barrier.
accumulator_pipeline_params.producer_arv_count = 1;
accumulator_pipeline_params.consumer_arv_count = size(AtomThrShapeMNK{}) * NumEpilogueThreads;
accumulator_pipeline_params.initializing_warp = 5;
AccumulatorPipeline accumulator_pipeline(shared_storage.pipelines.accumulator,
accumulator_pipeline_params,
cluster_shape,
cute::true_type{}, // Perform barrier init
cute::false_type{}); // Delay mask calculation
// CLC throttle pipeline
typename CLCThrottlePipeline::Params clc_throttle_pipeline_params;
if (WarpCategory::MainloopLoad == warp_category) {
clc_throttle_pipeline_params.role = CLCThrottlePipeline::ThreadCategory::Producer;
}
if (WarpCategory::Sched == warp_category) {
clc_throttle_pipeline_params.role = CLCThrottlePipeline::ThreadCategory::Consumer;
}
clc_throttle_pipeline_params.producer_arv_count = NumMainloopLoadThreads;
clc_throttle_pipeline_params.consumer_arv_count = NumSchedThreads;
clc_throttle_pipeline_params.dst_blockid = 0;
clc_throttle_pipeline_params.initializing_warp = 3;
CLCThrottlePipeline clc_throttle_pipeline(shared_storage.pipelines.clc_throttle, clc_throttle_pipeline_params);
CLCThrottlePipelineState clc_pipe_throttle_consumer_state;
CLCThrottlePipelineState clc_pipe_throttle_producer_state = cutlass::make_producer_start_state<CLCThrottlePipeline>();
// Tmem allocator
TmemAllocator tmem_allocator{};
// Sync allocation status between MMA and epilogue warps within CTA
cutlass::arch::NamedBarrier tmem_allocation_result_barrier(NumMMAThreads + NumEpilogueThreads, cutlass::arch::ReservedNamedBarriers::TmemAllocBarrier);
// Sync deallocation status between MMA warps of peer CTAs
cutlass::arch::ClusterBarrier& tmem_deallocation_result_barrier = shared_storage.pipelines.tmem_dealloc;
[[maybe_unused]] uint32_t dealloc_barrier_phase = 0;
if (WarpCategory::MMA == warp_category) {
if constexpr(!IsOverlappingAccum) {
if (has_mma_peer_cta && lane_predicate) {
tmem_deallocation_result_barrier.init(NumMMAThreads);
}
}
else {
if (has_mma_peer_cta && lane_predicate) {
tmem_deallocation_result_barrier.init(NumEpilogueThreads*2);
}
else if (lane_predicate) {
tmem_deallocation_result_barrier.init(NumEpilogueThreads);
}
}
}
// We need this to guarantee that the Pipeline init is visible
// To all producers and consumer threadblocks in the cluster
cutlass::pipeline_init_arrive_relaxed(cluster_size);
auto load_inputs = collective_mainloop.load_init(
problem_shape_MNKL, shared_storage.tensors.mainloop);
MainloopPipelineState mainloop_pipe_consumer_state;
MainloopPipelineState mainloop_pipe_producer_state = cutlass::make_producer_start_state<MainloopPipeline>();
EpiLoadPipelineState epi_load_pipe_consumer_state;
EpiLoadPipelineState epi_load_pipe_producer_state = cutlass::make_producer_start_state<EpiLoadPipeline>();
// epilogue store pipe is producer-only (consumer is TMA unit, waits via scoreboarding)
EpiStorePipelineState epi_store_pipe_producer_state = cutlass::make_producer_start_state<EpiStorePipeline>();
CLCPipelineState clc_pipe_consumer_state;
CLCPipelineState clc_pipe_producer_state = cutlass::make_producer_start_state<CLCPipeline>();
AccumulatorPipelineState accumulator_pipe_consumer_state;
AccumulatorPipelineState accumulator_pipe_producer_state = cutlass::make_producer_start_state<AccumulatorPipeline>();
dim3 block_id_in_cluster = cute::block_id_in_cluster();
// Calculate mask after cluster barrier arrival
mainloop_pipeline.init_masks(cluster_shape, block_id_in_cluster);
accumulator_pipeline.init_masks(cluster_shape, block_id_in_cluster);
// TileID scheduler
TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, block_id_in_cluster);
typename TileScheduler::WorkTileInfo work_tile_info = scheduler.initial_work_tile_info(cluster_shape);
auto cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info);
//
// TMEM "Allocation"
//
auto tmem_storage = collective_mainloop.template init_tmem_tensors<EpilogueTile, IsOverlappingAccum>(EpilogueTile{});
cutlass::pipeline_init_wait(cluster_size);
if (is_participant.main_load) {
// Ensure that the prefetched kernel does not touch
// unflushed global memory prior to this instruction
cutlass::arch::wait_on_dependent_grids();
bool do_load_order_arrive = is_epi_load_needed;
bool requires_clc_query = true;
do {
// Get the number of K tiles to compute for this work as well as the starting K tile offset of the work.
auto k_tile_iter = scheduler.get_k_tile_iterator(work_tile_info, problem_shape_MNKL, CtaShape_MNK{}, load_inputs.k_tiles);
auto k_tile_count = TileScheduler::get_work_k_tile_count(work_tile_info, problem_shape_MNKL, CtaShape_MNK{});
auto k_tile_prologue = min(MainloopPipeline::Stages, k_tile_count);
if constexpr (IsSchedDynamicPersistent) {
if (is_first_cta_in_cluster && requires_clc_query) {
clc_throttle_pipeline.producer_acquire(clc_pipe_throttle_producer_state);
clc_throttle_pipeline.producer_commit(clc_pipe_throttle_producer_state);
++clc_pipe_throttle_producer_state;
}
}
// Start mainloop prologue loads, arrive on the epilogue residual load barrier, resume mainloop loads
auto [mainloop_producer_state_next, k_tile_iter_next] = collective_mainloop.load(
mainloop_pipeline,
mainloop_pipe_producer_state,
load_inputs,
cta_coord_mnkl,
k_tile_iter, k_tile_prologue
);
mainloop_pipe_producer_state = mainloop_producer_state_next;
if (do_load_order_arrive) {
load_order_barrier.arrive();
do_load_order_arrive = false;
}
auto [mainloop_producer_state_next_, unused_] = collective_mainloop.load(
mainloop_pipeline,
mainloop_pipe_producer_state,
load_inputs,
cta_coord_mnkl,
k_tile_iter_next, k_tile_count - k_tile_prologue
);
mainloop_pipe_producer_state = mainloop_producer_state_next_;
// Sync warp to prevent non-participating threads entering next wave early
__syncwarp();
auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work(
work_tile_info,
clc_pipeline,
clc_pipe_consumer_state
);
work_tile_info = next_work_tile_info;
cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info);
requires_clc_query = increment_pipe;
if (increment_pipe) {
++clc_pipe_consumer_state;
}
} while (work_tile_info.is_valid());
collective_mainloop.load_tail(mainloop_pipeline, mainloop_pipe_producer_state);
}
else if (is_participant.sched) {
if constexpr (IsSchedDynamicPersistent) {
// Whether a new CLC query must be performed.
// See comment below where this variable is updated for a description of
// why this variable is needed.
bool requires_clc_query = true;
cutlass::arch::wait_on_dependent_grids();
do {
if (requires_clc_query) {
// Throttle CLC query to mitigate workload imbalance caused by skews among persistent workers.
clc_throttle_pipeline.consumer_wait(clc_pipe_throttle_consumer_state);
clc_throttle_pipeline.consumer_release(clc_pipe_throttle_consumer_state);
++clc_pipe_throttle_consumer_state;
// Query next clcID and update producer state
clc_pipe_producer_state = scheduler.advance_to_next_work(clc_pipeline, clc_pipe_producer_state);
}
// Fetch next work tile
auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work(
work_tile_info,
clc_pipeline,
clc_pipe_consumer_state
);
// Only perform a new CLC query if we consumed a new CLC query result in
// `fetch_next_work`. An example of a case in which CLC `fetch_next_work` does
// not consume a new CLC query response is when processing strxxm-K units.
// The current strxxm-K scheduler uses single WorkTileInfo to track multiple
// (potentially-partial) tiles to be computed via strxxm-K. In this case,
// `fetch_next_work` simply performs in-place updates on the existing WorkTileInfo,
// rather than consuming a CLC query response.
requires_clc_query = increment_pipe;
if (increment_pipe) {
++clc_pipe_consumer_state;
}
work_tile_info = next_work_tile_info;
} while (work_tile_info.is_valid());
clc_pipeline.producer_tail(clc_pipe_producer_state);
}
}
else if (is_participant.mma) {
// Tmem allocation sequence
tmem_allocator.allocate(TmemAllocator::Sm100TmemCapacityColumns, &shared_storage.tmem_base_ptr);
__syncwarp();
tmem_allocation_result_barrier.arrive();
uint32_t tmem_base_ptr = shared_storage.tmem_base_ptr;
collective_mainloop.set_tmem_offsets(tmem_storage, tmem_base_ptr);
auto mma_inputs = collective_mainloop.mma_init(
tmem_storage,
shared_storage.tensors.mainloop);
do {
auto k_tile_count = TileScheduler::get_work_k_tile_count(work_tile_info, problem_shape_MNKL, CtaShape_MNK{});
// Fetch next work tile
auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work(
work_tile_info,
clc_pipeline,
clc_pipe_consumer_state
);
if (increment_pipe) {
++clc_pipe_consumer_state;
}
// Accumulator stage slice
int acc_stage = [&] () {
if constexpr (IsOverlappingAccum) {
return accumulator_pipe_producer_state.phase() ^ 1;
}
else {
return accumulator_pipe_producer_state.index();
}
}();
if (is_mma_leader_cta) {
mainloop_pipe_consumer_state = collective_mainloop.mma(
cute::make_tuple(mainloop_pipeline, accumulator_pipeline),
cute::make_tuple(mainloop_pipe_consumer_state, accumulator_pipe_producer_state),
collective_mainloop.slice_accumulator(tmem_storage, acc_stage),
mma_inputs,
cta_coord_mnkl,
k_tile_count
);
accumulator_pipeline.producer_commit(accumulator_pipe_producer_state);
}
++accumulator_pipe_producer_state;
work_tile_info = next_work_tile_info;
cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info);
} while (work_tile_info.is_valid());
// Hint on an early release of global memory resources.
// The timing of calling this function only influences performance,
// not functional correctness.
cutlass::arch::launch_dependent_grids();
// Release the right to allocate before deallocations so that the next CTA can rasterize
tmem_allocator.release_allocation_lock();
if constexpr (!IsOverlappingAccum) {
// Leader MMA waits for leader + peer epilogues to release accumulator stage
if (is_mma_leader_cta) {
accumulator_pipeline.producer_tail(accumulator_pipe_producer_state);
}
// Signal to peer MMA that entire tmem allocation can be deallocated
if constexpr (has_mma_peer_cta) {
// Leader does wait + arrive, follower does arrive + wait
tmem_deallocation_result_barrier.arrive(mma_peer_cta_rank, not is_mma_leader_cta);
tmem_deallocation_result_barrier.wait(dealloc_barrier_phase);
tmem_deallocation_result_barrier.arrive(mma_peer_cta_rank, is_mma_leader_cta);
}
}
else {
tmem_deallocation_result_barrier.wait(dealloc_barrier_phase);
}
// Free entire tmem allocation
tmem_allocator.free(tmem_base_ptr, TmemAllocator::Sm100TmemCapacityColumns);
}
else if (is_participant.epi_load) {
// Ensure that the prefetched kernel does not touch
// unflushed global memory prior to this instruction
cutlass::arch::wait_on_dependent_grids();
bool do_load_order_wait = true;
bool do_tail_load = false;
int current_wave = 0;
do {
bool compute_epilogue = TileScheduler::compute_epilogue(work_tile_info, params.scheduler);
// Get current work tile and fetch next work tile
auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work(
work_tile_info,
clc_pipeline,
clc_pipe_consumer_state
);
work_tile_info = next_work_tile_info;
if (increment_pipe) {
++clc_pipe_consumer_state;
}
if (compute_epilogue) {
if (do_load_order_wait) {
load_order_barrier.wait();
do_load_order_wait = false;
}
bool reverse_epi_n = IsOverlappingAccum && (current_wave % 2 == 0);
epi_load_pipe_producer_state = collective_epilogue.template load<IsOverlappingAccum>(
epi_load_pipeline,
epi_load_pipe_producer_state,
problem_shape_MNKL,
CtaShape_MNK{},
cta_coord_mnkl,
TileShape{},
TiledMma{},
shared_storage.tensors.epilogue,
reverse_epi_n
);
do_tail_load = true;
}
current_wave++;
// Calculate the cta coordinates of the next work tile
cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info);
} while (work_tile_info.is_valid());
// Only perform a tail load if one of the work units processed performed
// an epilogue load. An example of a case in which a tail load should not be
// performed is in split-K if a cluster is only assigned non-final splits (for which
// the cluster does not compute the epilogue).
if (do_tail_load) {
collective_epilogue.load_tail(
epi_load_pipeline, epi_load_pipe_producer_state,
epi_store_pipeline, epi_store_pipe_producer_state);
}
}
else if (is_participant.epilogue) {
// Wait for tmem allocate here
tmem_allocation_result_barrier.arrive_and_wait();
uint32_t tmem_base_ptr = shared_storage.tmem_base_ptr;
collective_mainloop.set_tmem_offsets(tmem_storage, tmem_base_ptr);
bool do_tail_store = false;
do {
// Fetch next work tile
auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work(
work_tile_info,
clc_pipeline,
clc_pipe_consumer_state
);
if (increment_pipe) {
++clc_pipe_consumer_state;
}
// Accumulator stage slice
int acc_stage = [&] () {
if constexpr (IsOverlappingAccum) {
return accumulator_pipe_consumer_state.phase();
}
else {
return accumulator_pipe_consumer_state.index();
}
}();
auto accumulator = get<0>(collective_mainloop.slice_accumulator(tmem_storage, acc_stage));
accumulator_pipe_consumer_state = scheduler.template fixup<IsComplex>(
TiledMma{},
work_tile_info,
accumulator,
accumulator_pipeline,
accumulator_pipe_consumer_state,
typename CollectiveEpilogue::CopyOpT2R{}
);
//
// Epilogue and write to gD
//
if (scheduler.compute_epilogue(work_tile_info)) {
auto [load_state_next, store_state_next, acc_state_next] = collective_epilogue.template store<IsOverlappingAccum>(
epi_load_pipeline,
epi_load_pipe_consumer_state,
epi_store_pipeline,
epi_store_pipe_producer_state,
accumulator_pipeline,
accumulator_pipe_consumer_state,
problem_shape_MNKL,
CtaShape_MNK{},
cta_coord_mnkl,
TileShape{},
TiledMma{},
accumulator,
shared_storage.tensors.epilogue
);
epi_load_pipe_consumer_state = load_state_next;
epi_store_pipe_producer_state = store_state_next;
accumulator_pipe_consumer_state = acc_state_next;
do_tail_store = true;
}
work_tile_info = next_work_tile_info;
cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info);
} while (work_tile_info.is_valid());
if constexpr (IsOverlappingAccum) {
// Signal to peer MMA that Full TMEM alloc can be deallocated
if constexpr (has_mma_peer_cta) {
tmem_deallocation_result_barrier.arrive(mma_peer_cta_rank);
}
tmem_deallocation_result_barrier.arrive();
}
// Only perform a tail store if one of the work units processed performed
// an epilogue. An example of a case in which a tail load should not be
// performed is in split-K if a cluster is only assigned non-final splits (for which
// the cluster does not compute the epilogue).
if (do_tail_store) {
collective_epilogue.store_tail(
epi_load_pipeline, epi_load_pipe_consumer_state,
epi_store_pipeline, epi_store_pipe_producer_state,
CtaShape_MNK{});
}
}
else {
}
}
};
///////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::gemm::kernel
// ========== nvfp4_gemm.cu ==========
#include <ATen/cuda/CUDAContext.h>
#include <c10/cuda/CUDAGuard.h>
#include <cutlass/kernel_hardware_info.h>
#include <torch/all.h>
#include <cutlass/cutlass.h>
#include <cutlass/epilogue/thread/linear_combination.h>
#include <cutlass/gemm/dispatch_policy.hpp>
#include <cutlass/gemm/collective/collective_builder.hpp>
#include <cutlass/epilogue/collective/collective_builder.hpp>
#include <cutlass/detail/sm100_blockscaled_layout.hpp>
#include <cutlass/gemm/device/gemm_universal_adapter.h>
#include <cutlass/gemm/kernel/gemm_universal.hpp>
#include <cutlass/gemm/kernel/tile_scheduler_params.h>
#include <cute/tensor.hpp>
#include <cutlass/util/device_memory.h> // cutlass::device_memory::allocation
#include <cutlass/util/packed_stride.hpp>
#include <cutlass/gemm/kernel/tile_scheduler.hpp>
// Since we use load_inline everything monolithically. include "..." don't really matter.
using namespace cute;
// A matrix configuration
using ElementA = cutlass::nv_float4_t<cutlass::float_e2m1_t>; // Element type for A matrix operand
using ElementSFA = typename ElementA::ScaleFactorType; // Element type for SFA matrix operand
using LayoutATag = cutlass::layout::RowMajor; // Layout type for A matrix operand
constexpr int AlignmentA = 32; // Memory access granularity/alignment of A matrix in units of elements (up to 16 bytes)
// B matrix configuration
using ElementB = cutlass::nv_float4_t<cutlass::float_e2m1_t>; // Element type for A matrix operand
using ElementSFB = typename ElementB::ScaleFactorType; // Element type for SFA matrix operand
using LayoutBTag = cutlass::layout::ColumnMajor; // Layout type for B matrix operand
constexpr int AlignmentB = 32; // Memory access granularity/alignment of B matrix in units of elements (up to 16 bytes)
// C/D matrix configuration
using ElementC = cutlass::half_t; // Element type for C matrix operand
using LayoutCTag = cutlass::layout::RowMajor; // Layout type for C matrix operand
constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value; // Memory access granularity/alignment of C matrix in units of elements (up to 16 bytes)
// Kernel functional config
using ElementAccumulator = float; // Element type for internal accumulation
using ElementCompute = float; // Element type for internal accumulation
using ArchTag = cutlass::arch::Sm100; // Tag indicating the minimum SM that supports the intended feature
using OperatorClass = cutlass::arch::OpClassBlockScaledTensorOp; // Operator class tag
// Kernel Perf config
using MmaTileShape = Shape<_128,_128,_256>; // MMA's tile size
using ClusterShape = Shape<_1,_1,_1>; // Shape of the threadblocks in a cluster
constexpr int InputSFVectorSize = 16;
// C = alpha * acc
using FusionOperation = cutlass::epilogue::fusion::ScaledAcc<ElementC, ElementCompute>;
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag, OperatorClass,
MmaTileShape, ClusterShape,
cutlass::epilogue::collective::EpilogueTileAuto,
ElementAccumulator, ElementAccumulator,
ElementC, LayoutCTag, AlignmentC,
ElementC, LayoutCTag, AlignmentC,
cutlass::epilogue::collective::EpilogueScheduleAuto, // Epilogue schedule policy
FusionOperation
>::CollectiveOp;
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
ArchTag, OperatorClass,
ElementA, LayoutATag, AlignmentA,
ElementB, LayoutBTag, AlignmentB,
ElementAccumulator,
MmaTileShape, ClusterShape,
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
cutlass::gemm::collective::KernelScheduleAuto // Kernel schedule policy. Auto or using targeted scheduling policy
>::CollectiveOp;
using GemmKernel = gpu_mode::GemmUniversalCustom<
Shape<int,int,int, int>, // Indicates ProblemShape
CollectiveMainloop,
CollectiveEpilogue,
cutlass::gemm::DynamicPersistentScheduler>;
using Gemm = gpu_mode::GemmUniversalAdapterCustom<GemmKernel>;
// Reference device GEMM implementation type
using StrideA = typename Gemm::GemmKernel::StrideA;
using LayoutA = decltype(cute::make_layout(make_shape(0,0,0), StrideA{}));
using LayoutSFA = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFA; // Scale Factor tensors have an interleaved layout. Bring Layout instead of stride.
using StrideB = typename Gemm::GemmKernel::StrideB;
using LayoutB = decltype(cute::make_layout(make_shape(0,0,0), StrideB{}));
using LayoutSFB = typename Gemm::GemmKernel::CollectiveMainloop::LayoutSFB; // Scale Factor tensors have an interleaved layout. Bring Layout instead of stride.
using StrideC = typename Gemm::GemmKernel::StrideC;
using LayoutC = decltype(cute::make_layout(make_shape(0,0,0), StrideC{}));
using StrideD = typename Gemm::GemmKernel::StrideD;
using LayoutD = decltype(cute::make_layout(make_shape(0,0,0), StrideD{}));
using FusionOp = typename Gemm::EpilogueOutputOp;
at::Tensor nvfp4_gemm_blockscaled(
at::Tensor const &a,
at::Tensor const &b,
at::Tensor const &sfa_permuted,
at::Tensor const &sfb_permuted,
at::Tensor &c
) {
int const M = a.size(0);
int const N = b.size(0);
int const K_packed = a.size(1);
int const K = K_packed * 2;
using namespace cute;
// For SFA and SFB tensors layouts
using Sm1xxBlkScaledConfig = typename Gemm::GemmKernel::CollectiveMainloop::Sm1xxBlkScaledConfig;
StrideA a_stride = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, K, 1));
StrideB b_stride = cutlass::make_cute_packed_stride(StrideB{}, cute::make_shape(N, K, 1));
StrideC c_stride = cutlass::make_cute_packed_stride(StrideC{}, cute::make_shape(M, N, 1));
LayoutSFA layout_SFA = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(cute::make_shape(M, N, K, 1));
LayoutSFB layout_SFB = Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(cute::make_shape(M, N, K, 1));
typename Gemm::Arguments arguments {
cutlass::gemm::GemmUniversalMode::kGemm,
{M, N, K, 1},
{ // Mainloop arguments
static_cast<typename ElementA::DataType const*>(a.data_ptr()), a_stride,
static_cast<typename ElementB::DataType const*>(b.data_ptr()), b_stride,
static_cast<ElementSFA const*>(sfa_permuted.data_ptr()), layout_SFA,
static_cast<ElementSFB const*>(sfb_permuted.data_ptr()), layout_SFB
},
{ // Epilogue arguments
{ /* alpha = */ 1.f, /* beta = */ 0.f},
static_cast<ElementC*>(c.data_ptr()), c_stride,
static_cast<ElementC*>(c.data_ptr()), c_stride}
};
arguments.scheduler.max_swizzle_size = 0;
Gemm gemm;
size_t workspace_size = Gemm::get_workspace_size(arguments);
// Allocate workspace memory
cutlass::device_memory::allocation<uint8_t> workspace(workspace_size);
// Check if the problem size is supported or not
gemm.can_implement(arguments);
// Initialize CUTLASS kernel with arguments and workspace pointer
gemm.initialize(arguments, workspace.get());
gemm.run();
return c;
}
PYBIND11_MODULE(TORCH_EXTENSION_NAME, m) {
m.def("gemm_blockscaled", &nvfp4_gemm_blockscaled, "Nvfp4 Blockscaled GEMM");
}"""
nvfp4_module = load_inline(
name='gemm_blockscaled',
cpp_sources=[],
cuda_sources=[cuda_source],
extra_cuda_cflags=[
'-O3',
'-std=c++20',
'--expt-relaxed-constexpr',
'--expt-extended-lambda',
'--use_fast_math',
'-gencode=arch=compute_100a,code=sm_100a',
'-lineinfo',
'--ptxas-options=-v',
'-I/opt/cutlass/include',
'-I/usr/local/cuda/include',
],
verbose=False,
with_cuda=True
)
def custom_kernel(data: input_t) -> output_t:
"""Your custom kernel implementation."""
a, b, _, _, sfa_permuted, sfb_permuted, c = data
return nvfp4_module.gemm_blockscaled(a, b, sfa_permuted, sfb_permuted, c)scrolls · 1706 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