v4.5 dev update. (#3153)

This commit is contained in:
Junkai-Wu
2026-04-07 12:16:05 -04:00
committed by GitHub
parent 418d38a5de
commit a221da7ccf
265 changed files with 4913 additions and 1478 deletions
@@ -27,7 +27,7 @@
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
if (CUTLASS_NVCC_ARCHS MATCHES "120a|120f|121a")
if (CUTLASS_NVCC_ARCHS MATCHES "120a|121a")
cutlass_example_add_executable(
80a_blackwell_geforce_mxfp8_bf16_sparse_gemm
80a_blackwell_geforce_mxfp8_bf16_sparse_gemm.cu
@@ -0,0 +1,712 @@
/***************************************************************************************************
* Copyright (c) 2026 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: BSD-3-Clause
*
* Redistribution and use in source and binary forms, with or without
* modification, are permitted provided that the following conditions are met:
*
* 1. Redistributions of source code must retain the above copyright notice, this
* list of conditions and the following disclaimer.
*
* 2. Redistributions in binary form must reproduce the above copyright notice,
* this list of conditions and the following disclaimer in the documentation
* and/or other materials provided with the distribution.
*
* 3. Neither the name of the copyright holder nor the names of its
* contributors may be used to endorse or promote products derived from
* this software without specific prior written permission.
*
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
*
**************************************************************************************************/
/*
* Blackwell FP16 GEMM with Green Context
*
* The two schedulers share identical kernel configurations (element types, tile shapes,
* cluster shape, epilogue/mainloop builders) and differ only in:
*
* TileScheduler type:
* - DynamicPersistentScheduler (CLC-based)
* - StaticPersistentScheduler
*
* Use --scheduler=dynamic (default) or --scheduler=static to select the scheduler.
*/
#include <iostream>
#include <cuda.h>
#include "cutlass/cutlass.h"
#include "cute/tensor.hpp"
#include "cutlass/tensor_ref.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/gemm/device/gemm_universal_adapter.h"
#include "cutlass/gemm/kernel/gemm_universal.hpp"
#include "cutlass/gemm/kernel/tile_scheduler_params.h"
#include "cutlass/gemm/kernel/tile_scheduler.hpp"
#include "cutlass/util/command_line.h"
#include "cutlass/util/distribution.h"
#include "cutlass/util/host_tensor.h"
#include "cutlass/util/packed_stride.hpp"
#include "cutlass/util/tensor_view_io.h"
#include "cutlass/util/reference/device/gemm.h"
#include "cutlass/util/reference/device/tensor_compare.h"
#include "cutlass/util/reference/device/tensor_fill.h"
#include "helper.h"
/// Panic wrapper for CUDA Driver API errors (CUresult).
#define CUDA_DRIVER_CHECK(status) \
do { \
CUresult error = status; \
if (error != CUDA_SUCCESS) { \
const char *error_string; \
cuGetErrorString(error, &error_string); \
std::cerr << "Got CUDA driver error: " << error \
<< " (" << error_string << ")" \
<< " at: " << __LINE__ << std::endl; \
exit(EXIT_FAILURE); \
} \
} while (0)
/// Alias for CUDA Runtime API error checks (cudaError_t).
/// CUDA_CHECK is defined in helper.h; this alias makes the API boundary explicit.
#define CUDA_RUNTIME_CHECK CUDA_CHECK
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Common GEMM kernel configuration (shared by both schedulers)
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace gemm_common_config {
using namespace cute;
using ElementA = half_t;
using LayoutA = cutlass::layout::RowMajor;
constexpr int AlignmentA = 128 / cutlass::sizeof_bits<ElementA>::value;
using ElementB = half_t;
using LayoutB = cutlass::layout::ColumnMajor;
constexpr int AlignmentB = 128 / cutlass::sizeof_bits<ElementB>::value;
using ElementC = float;
using LayoutC = cutlass::layout::ColumnMajor;
constexpr int AlignmentC = 128 / cutlass::sizeof_bits<ElementC>::value;
using ElementAccumulator = float;
using ArchTag = cutlass::arch::Sm100;
using OperatorClass = cutlass::arch::OpClassTensorOp;
using MmaTileShape_MNK = Shape<_256,_128,_64>;
using ClusterShape_MNK = Shape<_2,_2,_1>;
using EpilogueScheduleType = cutlass::epilogue::TmaWarpSpecialized2Sm;
using KernelScheduleType = cutlass::gemm::KernelTmaWarpSpecialized2SmSm100;
} // namespace gemm_common_config
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Dynamic Persistent CLC kernel
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace dynamic_kernel_config {
using namespace cute;
using namespace gemm_common_config;
using TileScheduler = cutlass::gemm::DynamicPersistentScheduler;
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag, OperatorClass,
MmaTileShape_MNK, ClusterShape_MNK,
cutlass::epilogue::collective::EpilogueTileAuto,
ElementAccumulator, ElementAccumulator,
ElementC, LayoutC, AlignmentC,
ElementC, LayoutC, AlignmentC,
EpilogueScheduleType
>::CollectiveOp;
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
ArchTag, OperatorClass,
ElementA, LayoutA, AlignmentA,
ElementB, LayoutB, AlignmentB,
ElementAccumulator,
MmaTileShape_MNK, ClusterShape_MNK,
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
KernelScheduleType
>::CollectiveOp;
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
Shape<int,int,int, int>,
CollectiveMainloop,
CollectiveEpilogue,
TileScheduler>;
} // namespace dynamic_kernel_config
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Static Persistent kernel
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace static_kernel_config {
using namespace cute;
using namespace gemm_common_config;
using TileScheduler = cutlass::gemm::StaticPersistentScheduler;
using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder<
ArchTag, OperatorClass,
MmaTileShape_MNK, ClusterShape_MNK,
cutlass::epilogue::collective::EpilogueTileAuto,
ElementAccumulator, ElementAccumulator,
ElementC, LayoutC, AlignmentC,
ElementC, LayoutC, AlignmentC,
EpilogueScheduleType
>::CollectiveOp;
using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder<
ArchTag, OperatorClass,
ElementA, LayoutA, AlignmentA,
ElementB, LayoutB, AlignmentB,
ElementAccumulator,
MmaTileShape_MNK, ClusterShape_MNK,
cutlass::gemm::collective::StageCountAutoCarveout<static_cast<int>(sizeof(typename CollectiveEpilogue::SharedStorage))>,
KernelScheduleType
>::CollectiveOp;
using GemmKernel = cutlass::gemm::kernel::GemmUniversal<
Shape<int,int,int, int>,
CollectiveMainloop,
CollectiveEpilogue,
TileScheduler>;
} // namespace static_kernel_config
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Gemm device adapter types
/////////////////////////////////////////////////////////////////////////////////////////////////
using DynamicGemm = cutlass::gemm::device::GemmUniversalAdapter<dynamic_kernel_config::GemmKernel>;
using StaticGemm = cutlass::gemm::device::GemmUniversalAdapter<static_kernel_config::GemmKernel>;
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Common type aliases and global data (identical across both kernels)
/////////////////////////////////////////////////////////////////////////////////////////////////
using ElementA = gemm_common_config::ElementA;
using ElementB = gemm_common_config::ElementB;
using ElementC = gemm_common_config::ElementC;
using LayoutA = gemm_common_config::LayoutA;
using LayoutB = gemm_common_config::LayoutB;
using LayoutC = gemm_common_config::LayoutC;
using ElementAccumulator = gemm_common_config::ElementAccumulator;
using DeviceGemmReference = cutlass::reference::device::Gemm<
ElementA, LayoutA,
ElementB, LayoutB,
ElementC, LayoutC,
ElementAccumulator, ElementAccumulator>;
using StrideA = typename DynamicGemm::GemmKernel::StrideA;
using StrideB = typename DynamicGemm::GemmKernel::StrideB;
using StrideC = typename DynamicGemm::GemmKernel::StrideC;
using StrideD = typename DynamicGemm::GemmKernel::StrideD;
StrideA stride_A;
StrideB stride_B;
StrideC stride_C;
StrideD stride_D;
uint64_t seed = 0;
cutlass::DeviceAllocation<ElementA> block_A;
cutlass::DeviceAllocation<ElementB> block_B;
cutlass::DeviceAllocation<ElementC> block_C;
cutlass::DeviceAllocation<ElementC> block_D;
cutlass::DeviceAllocation<ElementC> block_ref_D;
#endif // defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
/// Testbed utility types
/////////////////////////////////////////////////////////////////////////////////////////////////
struct Options {
bool help;
bool use_cuda_graph;
int m, n, k;
float alpha, beta;
int iterations;
int swizzle;
int max_num_sm;
std::string raster_order;
std::string scheduler;
Options():
help(false),
use_cuda_graph(false),
m(8192), n(8192), k(8192),
alpha(1.f), beta(0.f),
iterations(30),
swizzle(0),
max_num_sm(0),
raster_order("heuristic"),
scheduler("dynamic")
{ }
void parse(int argc, char const **args) {
cutlass::CommandLine cmd(argc, args);
if (cmd.check_cmd_line_flag("help")) {
help = true;
return;
}
cmd.get_cmd_line_argument("m", m);
cmd.get_cmd_line_argument("n", n);
cmd.get_cmd_line_argument("k", k);
cmd.get_cmd_line_argument("alpha", alpha, 1.f);
cmd.get_cmd_line_argument("beta", beta, 0.f);
cmd.get_cmd_line_argument("iterations", iterations);
cmd.get_cmd_line_argument("swizzle", swizzle);
cmd.get_cmd_line_argument("max_num_sm", max_num_sm);
cmd.get_cmd_line_argument("raster_order", raster_order, std::string("heuristic"));
cmd.get_cmd_line_argument("scheduler", scheduler, std::string("dynamic"));
use_cuda_graph = cmd.check_cmd_line_flag("use_cuda_graph");
}
using RasterOrderOptions = cutlass::gemm::kernel::detail::RasterOrderOptions;
RasterOrderOptions get_raster_order() const {
if (raster_order == "along_m") {
return RasterOrderOptions::AlongM;
} else if (raster_order == "along_n") {
return RasterOrderOptions::AlongN;
} else {
return RasterOrderOptions::Heuristic;
}
}
std::ostream & print_usage(std::ostream &out) const {
out << "95_blackwell_gemm_green_context\n\n"
<< " Blackwell FP16 GEMM with Green Context support.\n"
<< " Supports both Dynamic Persistent CLC and Static Persistent schedulers.\n\n"
<< "Options:\n\n"
<< " --help If specified, displays this usage statement\n\n"
<< " --scheduler=<string> Tile scheduler: 'dynamic' (default) or 'static'\n"
<< " dynamic: DynamicPersistentScheduler (CLC-based)\n"
<< " static: StaticPersistentScheduler\n\n"
<< " --m=<int> Sets the M extent of the GEMM\n"
<< " --n=<int> Sets the N extent of the GEMM\n"
<< " --k=<int> Sets the K extent of the GEMM\n"
<< " --alpha=<f32> Epilogue scalar alpha\n"
<< " --beta=<f32> Epilogue scalar beta\n\n"
<< " --swizzle=<int> Cluster rasterization swizzle\n"
<< " --raster_order=<string> Raster order: 'heuristic' (default), 'along_m', or 'along_n'\n"
<< " --max_num_sm=<int> Max number of SMs for green context partition (0 = use all SMs, no green context)\n"
<< " --use_cuda_graph If specified, use CUDA graph capture/replay for profiling iterations\n\n"
<< " --iterations=<int> Number of profiling iterations to perform.\n\n";
out
<< "\n\nExamples:\n\n"
<< " # Dynamic scheduler, all SMs (no green context)\n"
<< " $ 95_blackwell_gemm_green_context --scheduler=dynamic --m=8192 --n=8192 --k=8192\n\n"
<< " # Static scheduler, 120-SM green context partition\n"
<< " $ 95_blackwell_gemm_green_context --scheduler=static --m=8192 --n=8192 --k=8192 --max_num_sm=120\n\n";
return out;
}
double gflops(double runtime_s) const
{
uint64_t flop = uint64_t(2) * m * n * k;
double gflop = double(flop) / double(1.0e9);
return gflop / runtime_s;
}
};
struct Result
{
double avg_runtime_ms;
double gflops;
cutlass::Status status;
cudaError_t error;
bool passed;
Result(
double avg_runtime_ms = 0,
double gflops = 0,
cutlass::Status status = cutlass::Status::kSuccess,
cudaError_t error = cudaSuccess)
:
avg_runtime_ms(avg_runtime_ms), gflops(gflops), status(status), error(error), passed(false)
{}
};
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
/////////////////////////////////////////////////////////////////////////////////////////////////
/// GEMM setup and evaluation
/////////////////////////////////////////////////////////////////////////////////////////////////
template <class Element>
void initialize_block(
cutlass::DeviceAllocation<Element>& block,
uint64_t seed=2023) {
Element scope_max, scope_min;
int bits_input = cutlass::sizeof_bits<Element>::value;
if (bits_input == 1) {
scope_max = Element(2);
scope_min = Element(0);
} else if (bits_input <= 8) {
scope_max = Element(2);
scope_min = Element(-2);
} else {
scope_max = Element(8);
scope_min = Element(-8);
}
cutlass::reference::device::BlockFillRandomUniform(
block.get(), block.size(), seed, scope_max, scope_min, 0);
}
void initialize(const Options &options) {
stride_A = cutlass::make_cute_packed_stride(StrideA{}, {options.m, options.k, 1});
stride_B = cutlass::make_cute_packed_stride(StrideB{}, {options.n, options.k, 1});
stride_C = cutlass::make_cute_packed_stride(StrideC{}, {options.m, options.n, 1});
stride_D = cutlass::make_cute_packed_stride(StrideD{}, {options.m, options.n, 1});
block_A.reset(options.m * options.k);
block_B.reset(options.k * options.n);
block_C.reset(options.m * options.n);
block_D.reset(options.m * options.n);
block_ref_D.reset(options.m * options.n);
initialize_block(block_A, seed + 2023);
initialize_block(block_B, seed + 2022);
initialize_block(block_C, seed + 2021);
}
/// Populates a Gemm::Arguments structure from the given commandline options.
///
/// The static persistent scheduler needs hw_info (sm_count, max_active_clusters) to compute
/// the launch grid size. When overlapping with another kernel via green context partitioning,
/// the green context stream MUST be passed so that cudaOccupancyMaxActiveClusters returns the
/// partition-scoped value. The dynamic persistent scheduler does not use hw_info.
///
/// We always set hw_info here for uniformity; it is harmless for the dynamic scheduler.
template <typename Gemm>
typename Gemm::Arguments args_from_options(const Options &options, cudaStream_t stream)
{
typename Gemm::Arguments arguments{
cutlass::gemm::GemmUniversalMode::kGemm,
{options.m, options.n, options.k, 1},
{block_A.get(), stride_A, block_B.get(), stride_B},
{{options.alpha, options.beta}, block_C.get(), stride_C, block_D.get(), stride_D}
};
arguments.scheduler.max_swizzle_size = options.swizzle;
arguments.scheduler.raster_order = options.get_raster_order();
// =====================================================================================
// IMPORTANT: Static persistent scheduler needs hw_info (sm_count, max_active_clusters)
// to compute the correct launch grid size.
//
// When the static persistent kernel overlaps with other kernel (e.g. via
// green context partitioning), you MUST pass the green context stream to
// make_kernel_hardware_info() so that cudaOccupancyMaxActiveClusters queries the max
// active clusters for that SM partition -- NOT the full device.
//
// If the stream is not passed, the query returns the full-device max active clusters,
// which leads to an oversized persistent grid that exceeds the partition's capacity,
// causing performance issues (because static scheduler stride to next work with launch grid size)
// =====================================================================================
arguments.hw_info = cutlass::KernelHardwareInfo::make_kernel_hardware_info<typename Gemm::GemmKernel>(
0 /* device_id */, 0 /* sm_count: auto-query */, 0 /* max_active_clusters: auto-query */, stream);
return arguments;
}
template <typename Gemm>
bool verify(const Options &options) {
cutlass::TensorRef ref_A(block_A.get(), LayoutA::packed({options.m, options.k}));
cutlass::TensorRef ref_B(block_B.get(), LayoutB::packed({options.k, options.n}));
cutlass::TensorRef ref_C(block_C.get(), LayoutC::packed({options.m, options.n}));
cutlass::TensorRef ref_D(block_ref_D.get(), LayoutC::packed({options.m, options.n}));
DeviceGemmReference gemm_reference;
gemm_reference(
{options.m, options.n, options.k},
ElementAccumulator(options.alpha),
ref_A,
ref_B,
ElementAccumulator(options.beta),
ref_C,
ref_D);
CUDA_RUNTIME_CHECK(cudaDeviceSynchronize());
bool passed = cutlass::reference::device::BlockCompareEqual(block_ref_D.get(), block_D.get(), block_D.size());
return passed;
}
/// Execute a given example GEMM computation.
/// Stream is REQUIRED -- pass cudaStreamDefault (0) for the full device, or a green context
/// stream for a partitioned launch.
template <typename Gemm>
int run(const Options &options, cudaStream_t stream)
{
initialize(options);
Gemm gemm;
auto arguments = args_from_options<Gemm>(options, stream);
dim3 grid = Gemm::get_grid_shape(arguments);
std::cout << " hw_info: sm_count=" << arguments.hw_info.sm_count
<< ", max_active_clusters=" << arguments.hw_info.max_active_clusters << std::endl;
std::cout << " Launch grid: (" << grid.x << ", " << grid.y << ", " << grid.z << ")" << std::endl;
size_t workspace_size = Gemm::get_workspace_size(arguments);
cutlass::device_memory::allocation<uint8_t> workspace(workspace_size);
CUTLASS_CHECK(gemm.can_implement(arguments));
CUTLASS_CHECK(gemm.initialize(arguments, workspace.get()));
CUTLASS_CHECK(gemm.run(stream));
CUDA_RUNTIME_CHECK(cudaStreamSynchronize(stream));
Result result;
result.passed = verify<Gemm>(options);
std::cout << " Disposition: " << (result.passed ? "Passed" : "Failed") << std::endl;
if (!result.passed) {
exit(-1);
}
if (options.iterations > 0)
{
GpuTimer timer;
if (options.use_cuda_graph) {
// cudaStreamBeginCapture cannot capture on the legacy default stream (stream 0).
cudaStream_t capture_stream = (stream == cudaStreamDefault) ? cudaStreamPerThread : stream;
CUTLASS_CHECK(gemm.initialize(arguments, workspace.get()));
cudaGraph_t graph;
cudaGraphExec_t graph_exec;
CUDA_RUNTIME_CHECK(cudaStreamBeginCapture(capture_stream, cudaStreamCaptureModeGlobal));
CUTLASS_CHECK(gemm.run(capture_stream));
CUDA_RUNTIME_CHECK(cudaStreamEndCapture(capture_stream, &graph));
CUDA_RUNTIME_CHECK(cudaGraphInstantiate(&graph_exec, graph, nullptr, nullptr, 0));
std::cout << " Using CUDA Graph for " << options.iterations << " iterations" << std::endl;
timer.start(stream);
for (int iter = 0; iter < options.iterations; ++iter) {
CUDA_RUNTIME_CHECK(cudaGraphLaunch(graph_exec, stream));
}
timer.stop();
CUDA_RUNTIME_CHECK(cudaGraphExecDestroy(graph_exec));
CUDA_RUNTIME_CHECK(cudaGraphDestroy(graph));
}
else {
timer.start(stream);
for (int iter = 0; iter < options.iterations; ++iter) {
CUTLASS_CHECK(gemm.initialize(arguments, workspace.get()));
CUTLASS_CHECK(gemm.run(stream));
}
timer.stop();
}
float elapsed_ms = timer.elapsed_millis();
result.avg_runtime_ms = double(elapsed_ms) / double(options.iterations);
result.gflops = options.gflops(result.avg_runtime_ms / 1000.0);
std::cout << " Problem Size: " << options.m << 'x' << options.n << 'x' << options.k << std::endl;
std::cout << " Avg runtime: " << result.avg_runtime_ms << " ms" << std::endl;
std::cout << " GFLOPS: " << result.gflops << std::endl;
}
return 0;
}
/// Dispatches to the correct Gemm type and handles green context setup.
template <typename Gemm>
int dispatch(const Options &options, int current_device_id)
{
if (options.max_num_sm > 0) {
//
// Green Context path: partition SMs and launch kernel on primary partition
//
CUDA_DRIVER_CHECK(cuInit(0));
CUdevice cu_device;
CUcontext primary_context;
CUDA_DRIVER_CHECK(cuDeviceGet(&cu_device, current_device_id));
CUDA_DRIVER_CHECK(cuDevicePrimaryCtxRetain(&primary_context, cu_device));
CUDA_DRIVER_CHECK(cuCtxSetCurrent(primary_context));
CUdevResource device_resource;
CUDA_DRIVER_CHECK(cuDeviceGetDevResource(cu_device, &device_resource, CU_DEV_RESOURCE_TYPE_SM));
std::cout << " Device SM count: " << device_resource.sm.smCount << std::endl;
if (options.max_num_sm >= static_cast<int>(device_resource.sm.smCount)) {
std::cout << " --max_num_sm (" << options.max_num_sm
<< ") >= device SM count (" << device_resource.sm.smCount
<< "), no green context split, using a dedicated stream with all SMs" << std::endl;
cudaStream_t stream;
CUDA_RUNTIME_CHECK(cudaStreamCreateWithFlags(&stream, cudaStreamNonBlocking));
int rc = run<Gemm>(options, stream);
CUDA_RUNTIME_CHECK(cudaStreamDestroy(stream));
CUDA_DRIVER_CHECK(cuDevicePrimaryCtxRelease(cu_device));
return rc;
}
CUdevResource primary_partition_resource;
CUdevResource remaining_partition_resource;
unsigned int num_groups = 1;
#if CUDA_VERSION >= 13000
unsigned int sm_alignment = device_resource.sm.smCoscheduledAlignment;
unsigned int aligned_sm = (static_cast<unsigned int>(options.max_num_sm) / sm_alignment) * sm_alignment;
if (aligned_sm == 0) {
aligned_sm = sm_alignment;
}
std::cout << " SM coscheduled alignment: " << sm_alignment << std::endl;
std::cout << " Requested --max_num_sm: " << options.max_num_sm
<< ", aligned (round down): " << aligned_sm << std::endl;
CUDA_DRIVER_CHECK(cuDevSmResourceSplitByCount(
&primary_partition_resource,
&num_groups,
&device_resource,
&remaining_partition_resource,
0,
aligned_sm));
#else
std::cout << " Requested --max_num_sm: " << options.max_num_sm << std::endl;
CUDA_DRIVER_CHECK(cuDevSmResourceSplitByCount(
&primary_partition_resource,
&num_groups,
&device_resource,
&remaining_partition_resource,
0,
static_cast<unsigned int>(options.max_num_sm)));
#endif
std::cout << " Primary partition SM count (for GEMM): " << primary_partition_resource.sm.smCount << std::endl;
std::cout << " Remaining partition SM count: " << remaining_partition_resource.sm.smCount << std::endl;
CUdevResourceDesc primary_partition_desc;
CUdevResourceDesc remaining_partition_desc;
CUDA_DRIVER_CHECK(cuDevResourceGenerateDesc(&primary_partition_desc, &primary_partition_resource, 1));
CUDA_DRIVER_CHECK(cuDevResourceGenerateDesc(&remaining_partition_desc, &remaining_partition_resource, 1));
CUgreenCtx primary_partition_green_ctx;
CUgreenCtx remaining_partition_green_ctx;
CUDA_DRIVER_CHECK(cuGreenCtxCreate(&primary_partition_green_ctx, primary_partition_desc, cu_device, CU_GREEN_CTX_DEFAULT_STREAM));
CUDA_DRIVER_CHECK(cuGreenCtxCreate(&remaining_partition_green_ctx, remaining_partition_desc, cu_device, CU_GREEN_CTX_DEFAULT_STREAM));
// Remaining partition stream is not used in this example but instantiated to show how.
CUstream primary_partition_cu_stream;
CUstream remaining_partition_cu_stream;
CUDA_DRIVER_CHECK(cuGreenCtxStreamCreate(&primary_partition_cu_stream, primary_partition_green_ctx, CU_STREAM_NON_BLOCKING, 0));
CUDA_DRIVER_CHECK(cuGreenCtxStreamCreate(&remaining_partition_cu_stream, remaining_partition_green_ctx, CU_STREAM_NON_BLOCKING, 0));
cudaStream_t primary_partition_stream = static_cast<cudaStream_t>(primary_partition_cu_stream);
int rc = run<Gemm>(options, primary_partition_stream);
CUDA_DRIVER_CHECK(cuStreamDestroy(primary_partition_cu_stream));
CUDA_DRIVER_CHECK(cuStreamDestroy(remaining_partition_cu_stream));
CUDA_DRIVER_CHECK(cuGreenCtxDestroy(primary_partition_green_ctx));
CUDA_DRIVER_CHECK(cuGreenCtxDestroy(remaining_partition_green_ctx));
CUDA_DRIVER_CHECK(cuDevicePrimaryCtxRelease(cu_device));
return rc;
}
else {
return run<Gemm>(options, cudaStreamDefault);
}
}
#endif // defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
///////////////////////////////////////////////////////////////////////////////////////////////////
int main(int argc, char const **args) {
if (__CUDACC_VER_MAJOR__ < 12 || (__CUDACC_VER_MAJOR__ == 12 && __CUDACC_VER_MINOR__ < 9)) {
std::cerr << "This example requires CUDA 12.9 or newer." << std::endl;
return 0;
}
cudaDeviceProp props;
int current_device_id;
CUDA_RUNTIME_CHECK(cudaGetDevice(&current_device_id));
CUDA_RUNTIME_CHECK(cudaGetDeviceProperties(&props, current_device_id));
if (props.major != 10 || props.minor != 0) {
std::cerr << "This example requires a GPU with compute capability 100a." << std::endl;
return 0;
}
Options options;
options.parse(argc, args);
if (options.help) {
options.print_usage(std::cout) << std::endl;
return 0;
}
#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
if (options.scheduler == "static") {
std::cout << "Using StaticPersistentScheduler" << std::endl;
return dispatch<StaticGemm>(options, current_device_id);
}
else if (options.scheduler == "dynamic") {
std::cout << "Using DynamicPersistentScheduler" << std::endl;
return dispatch<DynamicGemm>(options, current_device_id);
}
else {
std::cerr << "Unknown scheduler: '" << options.scheduler
<< "'. Use --scheduler=dynamic or --scheduler=static." << std::endl;
return 1;
}
#endif // defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED)
return 0;
}
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -0,0 +1,54 @@
# Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
# SPDX-License-Identifier: BSD-3-Clause
#
# Redistribution and use in source and binary forms, with or without
# modification, are permitted provided that the following conditions are met:
#
# 1. Redistributions of source code must retain the above copyright notice, this
# list of conditions and the following disclaimer.
#
# 2. Redistributions in binary form must reproduce the above copyright notice,
# this list of conditions and the following disclaimer in the documentation
# and/or other materials provided with the distribution.
#
# 3. Neither the name of the copyright holder nor the names of its
# contributors may be used to endorse or promote products derived from
# this software without specific prior written permission.
#
# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
set(TEST_DYNAMIC_MAX_NUM_SM_64 --scheduler=dynamic --max_num_sm=64)
set(TEST_DYNAMIC_MAX_NUM_SM_120 --scheduler=dynamic --max_num_sm=120)
set(TEST_DYNAMIC_MAX_NUM_SM_148 --scheduler=dynamic --max_num_sm=148)
set(TEST_DYNAMIC_MAX_NUM_SM_160 --scheduler=dynamic --max_num_sm=160)
set(TEST_STATIC_MAX_NUM_SM_64 --scheduler=static --max_num_sm=64)
set(TEST_STATIC_MAX_NUM_SM_120 --scheduler=static --max_num_sm=120)
set(TEST_STATIC_MAX_NUM_SM_148 --scheduler=static --max_num_sm=148)
set(TEST_STATIC_MAX_NUM_SM_160 --scheduler=static --max_num_sm=160)
if(CUTLASS_NVCC_ARCHS MATCHES "100a|100f|101a|101f|103a|103f")
cutlass_example_add_executable(
95_blackwell_gemm_green_context
95_blackwell_gemm_green_context.cu
TEST_COMMAND_OPTIONS
TEST_DYNAMIC_MAX_NUM_SM_64
TEST_DYNAMIC_MAX_NUM_SM_120
TEST_DYNAMIC_MAX_NUM_SM_148
TEST_DYNAMIC_MAX_NUM_SM_160
TEST_STATIC_MAX_NUM_SM_64
TEST_STATIC_MAX_NUM_SM_120
TEST_STATIC_MAX_NUM_SM_148
TEST_STATIC_MAX_NUM_SM_160
)
endif()
@@ -0,0 +1,114 @@
# Example 95: Blackwell GEMM with Green Context
[Green Context](https://docs.nvidia.com/cuda/cuda-driver-api/group__CUDA__GREEN__CONTEXTS.html) is a lightweight method to partition GPU SM resources.
This example demonstrates a Blackwell FP16 GEMM with green context support using two tile scheduler variants, selectable at runtime via `--scheduler`:
1. **Dynamic Persistent CLC** (`--scheduler=dynamic`, default): Uses `DynamicPersistentScheduler` (Cluster Launch Control based). No modification needed for green context -- simply launch kernel onto partition stream.
2. **Static Persistent** (`--scheduler=static`): Uses `StaticPersistentScheduler`. For green context, use partition stream to query max active cluster; modify launch grid based on partition stream max active cluster.
For SM90 Hopper dynamic / static kernels, modifying the launch grid is required.
## Build
From the CUTLASS build directory:
```shell
# Configure (only examples, SM100a)
cmake <cutlass_root> \
-DCUTLASS_NVCC_ARCHS=100a \
-DCUTLASS_ENABLE_EXAMPLES=ON \
-DCUTLASS_ENABLE_TESTS=OFF \
-DCUTLASS_ENABLE_LIBRARY=OFF \
-DCUTLASS_ENABLE_PROFILER=OFF \
-DCMAKE_BUILD_TYPE=Release
# Build only this example
make 95_blackwell_gemm_green_context -j$(nproc)
```
## Run
### Dynamic Persistent CLC scheduler
#### Without green context (all SMs)
```shell
./examples/95_blackwell_gemm_green_context/95_blackwell_gemm_green_context \
--scheduler=dynamic --m=8192 --n=8192 --k=8192 --iterations=30
```
#### With green context (partitioned SMs)
Use `--max_num_sm` to specify the number of SMs for the primary partition (GEMM workload).
The remaining SMs are assigned to the remaining partition.
```shell
./examples/95_blackwell_gemm_green_context/95_blackwell_gemm_green_context \
--scheduler=dynamic --m=8192 --n=8192 --k=8192 --max_num_sm=120 --iterations=30 --raster_order=along_m
./examples/95_blackwell_gemm_green_context/95_blackwell_gemm_green_context \
--scheduler=dynamic --m=8192 --n=8192 --k=8192 --max_num_sm=120 --iterations=30 --raster_order=along_n
```
### Static Persistent scheduler
#### Without green context (all SMs)
```shell
./examples/95_blackwell_gemm_green_context/95_blackwell_gemm_green_context \
--scheduler=static --m=8192 --n=8192 --k=8192 --iterations=30
```
#### With green context (partitioned SMs)
```shell
./examples/95_blackwell_gemm_green_context/95_blackwell_gemm_green_context \
--scheduler=static --m=8192 --n=8192 --k=8192 --max_num_sm=120 --iterations=30 --raster_order=along_m
./examples/95_blackwell_gemm_green_context/95_blackwell_gemm_green_context \
--scheduler=static --m=8192 --n=8192 --k=8192 --max_num_sm=120 --iterations=30 --raster_order=along_n
```
## Nsight Systems Profiling
Use `nsys profile` to capture the kernel execution under green context partitioning.
In the Nsys UI, you should see the green context section on the left panel, and the GEMM kernel
launched onto that green context partition.
### Profile Dynamic Persistent CLC scheduler
#### Without green context
```shell
nsys profile -o gemm_dynamic_all_sm \
./examples/95_blackwell_gemm_green_context/95_blackwell_gemm_green_context \
--scheduler=dynamic --m=8192 --n=8192 --k=8192 --iterations=30
```
#### With green context (120 SMs for GEMM)
```shell
nsys profile -o gemm_dynamic_green_ctx_120sm \
./examples/95_blackwell_gemm_green_context/95_blackwell_gemm_green_context \
--scheduler=dynamic --m=8192 --n=8192 --k=8192 --max_num_sm=120 --iterations=30
```
### Profile Static Persistent scheduler
#### Without green context
```shell
nsys profile -o gemm_static_all_sm \
./examples/95_blackwell_gemm_green_context/95_blackwell_gemm_green_context \
--scheduler=static --m=8192 --n=8192 --k=8192 --iterations=30
```
#### With green context (120 SMs for GEMM)
```shell
nsys profile -o gemm_static_green_ctx_120sm \
./examples/95_blackwell_gemm_green_context/95_blackwell_gemm_green_context \
--scheduler=static --m=8192 --n=8192 --k=8192 --max_num_sm=120 --iterations=30
```
+1
View File
@@ -172,6 +172,7 @@ foreach(EXAMPLE
92_blackwell_moe_gemm
93_blackwell_low_latency_gqa
94_ada_fp8_blockwise
95_blackwell_gemm_green_context
111_hopper_ssd
112_blackwell_ssd
)
@@ -26,9 +26,11 @@
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
import cutlass.cute as cute
import cutlass
"""
Example of automatic shared memory size computation for configuring kernel launch
@@ -51,11 +53,15 @@ class SharedData:
@cute.kernel
def kernel():
def kernel_static():
"""
Example kernel that allocates shared memory.
The total allocation will be automatically calculated when smem=None.
"""
tidx, _, _ = cute.arch.block_idx()
if tidx == 0:
cute.printf("Running kernel_static")
allocator = cutlass.utils.SmemAllocator()
# Allocate various types of shared memory
@@ -68,6 +74,8 @@ def kernel():
byte_alignment=16,
swizzle=None,
)
cute.printf("Kernel launch smem size: {}", cute.arch.dynamic_smem_size())
return
@@ -79,7 +87,7 @@ def kernel_no_smem():
"""
tidx, _, _ = cute.arch.block_idx()
if tidx == 0:
cute.printf("Hello world")
cute.printf("Running kernel_no_smem")
return
@@ -89,26 +97,49 @@ if __name__ == "__main__":
print("Launching kernel with auto smem size. (launch config `smem=None`)")
# Compile the example
# Compile the static example
@cute.jit
def launch_kernel1():
k = kernel()
k.launch(
def launch_kernelno_smem():
kernel_no_smem().launch(
grid=(1, 1, 1),
block=(1, 1, 1),
)
print(f"Kernel recorded internal smem usage: {k.smem_usage()}")
# --------
print(f" > Run {kernel_no_smem.__name__}")
func = cute.compile(launch_kernelno_smem)
func()
cutlass.cuda.stream_sync(cutlass.cuda.default_stream())
@cute.jit
def launch_kernel2():
k = kernel_no_smem()
k.launch(
def launch_kernel_static():
kernel_static().launch(
grid=(1, 1, 1),
block=(1, 1, 1),
# smem=None
# auto infer launch kernel static smem usage
)
print(f"Kernel recorded internal smem usage: {k.smem_usage()}")
cute.compile(launch_kernel1)
cute.compile(launch_kernel2)
# --------
print(f" > Run {kernel_static.__name__} with sufficient smem")
func = cute.compile(launch_kernel_static)
func()
cutlass.cuda.stream_sync(cutlass.cuda.default_stream())
@cute.jit
def launch_kernel_static_insufficient():
kernel_static().launch(
grid=(1, 1, 1),
block=(1, 1, 1),
# launch kernel with static smem usage exceeds cfg
# show warning
smem=16,
)
# --------
print(f" > Run {kernel_static.__name__} with insufficient smem, show warning:")
func = cute.compile(launch_kernel_static_insufficient)
func()
cutlass.cuda.stream_sync(cutlass.cuda.default_stream())
print("PASS")
@@ -265,7 +265,6 @@ class HSTUAttentionForwardAmpere(object):
).launch(
grid=grid_dim,
block=[self._num_threads, 1, 1],
smem=SharedStorage.size_in_bytes(),
stream=stream,
)
@@ -129,8 +129,8 @@ def kernel(
# ptr<i64, smem, align<128>>
# ptr<f32, smem, align<8>>
print(struct_in_smem.a.data_ptr())
print(struct_in_smem.b)
print(struct_in_smem.c.real)
print(struct_in_smem.b.ptr)
print(struct_in_smem.c.real.ptr)
# ptr<i8, smem, align<512>>
print(section_in_smem)
# ptr<i64, smem, align<64>>
@@ -138,6 +138,17 @@ def kernel(
# tensor<ptr<f16, smem, align<32>> o (16,4):(1,16)>
print(tensor_in_smem)
# assign struct member array element
cute.printf("struct_in_smem.a[0] = {}", struct_in_smem.a[0])
struct_in_smem.a[0] = 2
cute.printf("struct_in_smem.a[0] = {}", struct_in_smem.a[0])
# assign struct member scalar
cute.printf("struct_in_smem.b.ptr = {}", struct_in_smem.b.ptr)
cute.printf("struct_in_smem.b: value = {}", struct_in_smem.b.ptr.load())
struct_in_smem.b = 16
cute.printf("struct_in_smem.b: value = {}", struct_in_smem.b.ptr.load())
# fill MemRange tensor in struct and copy to dst
a_tensor = struct_in_smem.a.get_tensor(cute.make_layout((8, 4)))
a_tensor.fill(const_a)
@@ -169,7 +180,9 @@ def host(
):
# Note: Shared Memory size is automatically calculated now
kernel(const_a, dst_a, const_b, dst_b, const_c, dst_c).launch(
grid=(1, 1, 1), block=(1, 1, 1)
grid=(1, 1, 1),
block=(1, 1, 1),
# Automatically calculate the launch kernel shared memory usage when `smem=None`
)
+27 -15
View File
@@ -175,15 +175,6 @@ class TensorOpGemm:
(self.cta_tiler[0], self.cta_tiler[1]),
)
# Shared memory allocated for operations with A, B will be
# overwritten for operations on C. This is to improve performance
# by reducing the size of shared memory requested by each block
smem_size = max(
cute.size_in_bytes(mC.element_type, sC_layout),
cute.size_in_bytes(mA.element_type, sA_layout)
+ cute.size_in_bytes(mB.element_type, sB_layout),
)
# ///////////////////////////////////////////////////////////////////////////////
# Tiled copy:
# The majorness of tA/tB/tC follows the majorness of gA/gB/gC,
@@ -282,7 +273,6 @@ class TensorOpGemm:
).launch(
grid=rasterization_remap_grid_dim,
block=[self.num_threads, 1, 1],
smem=smem_size,
)
@cute.kernel
@@ -382,14 +372,36 @@ class TensorOpGemm:
# tAgA: (CPY, CPY_M, CPY_K, k) , tBgB: (CPY, CPY_N, CPY_K, k)
# tAsA: (CPY, CPY_M, CPY_K, PIPE) , tBsB: (CPY, CPY_N, CPY_K, PIPE)
# ///////////////////////////////////////////////////////////////////////////////
@cute.struct
class SharedStorageAB:
a: cute.struct.Align[
cute.struct.MemRange[mA.element_type, cute.cosize(sA_layout)],
16,
]
b: cute.struct.Align[
cute.struct.MemRange[mB.element_type, cute.cosize(sB_layout)],
16,
]
@cute.struct
class SharedStorageC:
c: cute.struct.Align[
cute.struct.MemRange[mC.element_type, cute.cosize(sC_layout)],
16,
]
# Shared memory buffer
smem = cutlass.utils.SmemAllocator()
sA = smem.allocate_tensor(mA.element_type, sA_layout, 16)
sB = smem.allocate_tensor(mB.element_type, sB_layout, 16)
sC = cute.make_tensor(
cute.recast_ptr(sA.iterator, dtype=self.c_dtype), sC_layout
# Shared memory allocated for operations with A, B will be
# overwritten for operations on C. This is to improve performance
# by reducing the size of shared memory requested by each block
storage = smem.allocate(
max(SharedStorageAB.size_in_bytes(), SharedStorageC.size_in_bytes()),
byte_alignment=16,
)
sA = SharedStorageAB(storage).a.get_tensor(sA_layout)
sB = SharedStorageAB(storage).b.get_tensor(sB_layout)
sC = SharedStorageC(storage).c.get_tensor(sC_layout)
thr_copy_A = tiled_copy_A.get_slice(tidx)
thr_copy_B = tiled_copy_B.get_slice(tidx)
@@ -549,7 +549,7 @@ class BlockwiseGemmKernel:
cutlass.Int64, self.num_tile_stage * 2
]
epi_mbar_ptr: cute.struct.MemRange[cutlass.Int64, 1 * 2]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# (EPI_TILE_M, EPI_TILE_N, STAGE)
sC: cute.struct.Align[
@@ -614,7 +614,6 @@ class BlockwiseGemmKernel:
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
smem=self.shared_storage.size_in_bytes(),
stream=stream,
min_blocks_per_mp=1,
)
@@ -682,9 +681,6 @@ class BlockwiseGemmKernel:
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
tmem_dealloc_mbar_ptr = storage.tmem_dealloc_mbar_ptr
tmem_holding_buf = storage.tmem_holding_buf
# Initialize mainloop ab_pipeline (barrier) and states
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
@@ -771,11 +767,11 @@ class BlockwiseGemmKernel:
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -568,7 +568,7 @@ class BlockwiseContiguousGroupedGemmKernel:
cutlass.Int64, self.num_tile_stage * 2
]
epi_mbar_ptr: cute.struct.MemRange[cutlass.Int64, 1 * 2]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# (EPI_TILE_M, EPI_TILE_N, STAGE)
sC: cute.struct.Align[
@@ -634,7 +634,6 @@ class BlockwiseContiguousGroupedGemmKernel:
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
smem=self.shared_storage.size_in_bytes(),
stream=stream,
min_blocks_per_mp=1,
)
@@ -703,9 +702,6 @@ class BlockwiseContiguousGroupedGemmKernel:
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
tmem_dealloc_mbar_ptr = storage.tmem_dealloc_mbar_ptr
tmem_holding_buf = storage.tmem_holding_buf
# Initialize mainloop ab_pipeline (barrier) and states
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
@@ -792,11 +788,11 @@ class BlockwiseContiguousGroupedGemmKernel:
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -567,7 +567,7 @@ class BlockwiseMaskedGroupedGemmKernel:
cutlass.Int64, self.num_tile_stage * 2
]
epi_mbar_ptr: cute.struct.MemRange[cutlass.Int64, 1 * 2]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# (EPI_TILE_M, EPI_TILE_N, STAGE)
sC: cute.struct.Align[
@@ -633,7 +633,6 @@ class BlockwiseMaskedGroupedGemmKernel:
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
smem=self.shared_storage.size_in_bytes(),
stream=stream,
min_blocks_per_mp=1,
)
@@ -702,9 +701,6 @@ class BlockwiseMaskedGroupedGemmKernel:
smem = utils.SmemAllocator()
storage = smem.allocate(self.shared_storage)
tmem_dealloc_mbar_ptr = storage.tmem_dealloc_mbar_ptr
tmem_holding_buf = storage.tmem_holding_buf
# Initialize mainloop ab_pipeline (barrier) and states
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1
@@ -791,11 +787,11 @@ class BlockwiseMaskedGroupedGemmKernel:
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -615,7 +615,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# (EPI_TILE_M, EPI_TILE_N, STAGE)
sC: cute.struct.Align[
@@ -794,11 +794,11 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -551,7 +551,7 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# (EPI_TILE_M, EPI_TILE_N, STAGE)
sC: cute.struct.Align[
@@ -737,11 +737,11 @@ class Sm100BlockScaledPersistentDenseGemmKernel:
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -517,7 +517,7 @@ class DenseGemmKernel:
acc_full_mbar_ptr: cute.struct.MemRange[
cutlass.Int64, self.num_acc_stage * 2
]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
smem = utils.SmemAllocator()
@@ -564,10 +564,10 @@ class DenseGemmKernel:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -493,7 +493,7 @@ class SM100PersistentDenseGemmAlphaBetaKernel:
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
c_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_c_stage]
c_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_c_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# (EPI_TILE_M, EPI_TILE_N, STAGE)
sD: cute.struct.Align[
@@ -674,11 +674,11 @@ class SM100PersistentDenseGemmAlphaBetaKernel:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_ids[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -593,7 +593,7 @@ class PersistentDenseGemmKernel:
acc_full_mbar_ptr: cute.struct.MemRange[
cutlass.Int64, self.num_acc_stage * 2
]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
smem = utils.SmemAllocator()
@@ -644,11 +644,11 @@ class PersistentDenseGemmKernel:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=self.epilogue_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -614,7 +614,7 @@ class PersistentDenseGemmKernel:
acc_full_mbar_ptr: cute.struct.MemRange[
cutlass.Int64, self.num_acc_stage * 2
]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
clc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, 2]
clc_response: cute.struct.MemRange[cutlass.Int32, 4]
@@ -686,11 +686,11 @@ class PersistentDenseGemmKernel:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=self.epilogue_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -514,7 +514,7 @@ class DenseGemmKernel:
acc_full_mbar_ptr: cute.struct.MemRange[
cutlass.Int64, self.num_acc_stage * 2
]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
smem = utils.SmemAllocator()
@@ -562,10 +562,10 @@ class DenseGemmKernel:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -573,7 +573,7 @@ class DenseGemmEFC:
# Barriers used by the supplemental load tensor pipeline.
c_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_c_stage]
c_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_c_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# (MMA, MMA_M, MMA_K, STAGE)
sA: cute.struct.Align[
@@ -651,11 +651,11 @@ class DenseGemmEFC:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=self.epilogue_warp_id[0],
is_two_cta=self.use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
+2 -2
View File
@@ -998,7 +998,7 @@ class BlackwellFusedMultiHeadAttentionForward:
# Alloc tmem buffer
tmem_alloc_cols = Int32(self.tmem_alloc_cols)
cute.arch.alloc_tmem(tmem_alloc_cols, storage.tmem_holding_buf)
cute.arch.alloc_tmem(tmem_alloc_cols, storage.tmem_holding_buf.ptr)
self.tmem_alloc_barrier.arrive_and_wait()
tile_sched = fmha_utils.create_fmha_static_tile_scheduler(
tile_sched_params, cute.arch.block_idx(), cute.arch.grid_dim()
@@ -1260,7 +1260,7 @@ class BlackwellFusedMultiHeadAttentionForward:
tmem_ptr = cute.arch.retrieve_tmem_ptr(
Float32,
alignment=16,
ptr_to_buffer_holding_addr=storage.tmem_holding_buf,
ptr_to_buffer_holding_addr=storage.tmem_holding_buf.ptr,
)
cute.arch.dealloc_tmem(tmem_ptr, tmem_alloc_cols)
@@ -708,7 +708,6 @@ class BlackwellFusedMultiHeadAttentionBackward:
grid=bwd_grid,
block=[self.threads_per_cta, 1, 1],
cluster=[1, 1, 1],
smem=self.shared_storage.size_in_bytes(),
stream=stream,
min_blocks_per_mp=1,
)
@@ -913,7 +912,7 @@ class BlackwellFusedMultiHeadAttentionBackward:
)
sLSE = storage.sLSE.get_tensor(LSE_smem_layout)
sSum_OdO = storage.sSum_OdO.get_tensor(sum_OdO_smem_layout)
tmem_holding_buf = storage.tmem_holding_buf
tmem_holding_buf = storage.tmem_holding_buf.ptr
sQT_ptr = cute.recast_ptr(sQ.iterator, QT_smem_layout_staged.inner)
sQT = cute.make_tensor(sQT_ptr, QT_smem_layout_staged.outer)
@@ -567,7 +567,7 @@ class Sm100GroupedBlockScaledGemmKernel:
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# (EPI_TILE_M, EPI_TILE_N, STAGE)
sC: cute.struct.Align[
@@ -641,7 +641,6 @@ class Sm100GroupedBlockScaledGemmKernel:
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=(*self.cluster_shape_mn, 1),
smem=self.shared_storage.size_in_bytes(),
stream=stream,
min_blocks_per_mp=1,
)
@@ -737,8 +736,8 @@ class Sm100GroupedBlockScaledGemmKernel:
+ Sm100GroupedBlockScaledGemmKernel.bytes_per_tensormap // 8
)
tmem_dealloc_mbar_ptr = storage.tmem_dealloc_mbar_ptr
tmem_holding_buf = storage.tmem_holding_buf
tmem_dealloc_mbar_ptr = storage.tmem_dealloc_mbar.ptr
tmem_holding_buf_ptr = storage.tmem_holding_buf.ptr
# Initialize mainloop ab_pipeline (barrier) and states
ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread)
@@ -1249,7 +1248,7 @@ class Sm100GroupedBlockScaledGemmKernel:
acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(
self.acc_dtype,
alignment=16,
ptr_to_buffer_holding_addr=tmem_holding_buf,
ptr_to_buffer_holding_addr=tmem_holding_buf_ptr,
)
# (MMA, MMA_M, MMA_N, STAGE)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
@@ -1446,7 +1445,7 @@ class Sm100GroupedBlockScaledGemmKernel:
if warp_idx == self.epilog_warp_id[0]:
cute.arch.alloc_tmem(
self.num_tmem_alloc_cols,
tmem_holding_buf,
tmem_holding_buf_ptr,
is_two_cta=use_2cta_instrs,
)
@@ -1461,7 +1460,7 @@ class Sm100GroupedBlockScaledGemmKernel:
acc_tmem_ptr = cute.arch.retrieve_tmem_ptr(
self.acc_dtype,
alignment=16,
ptr_to_buffer_holding_addr=tmem_holding_buf,
ptr_to_buffer_holding_addr=tmem_holding_buf_ptr,
)
# (MMA, MMA_M, MMA_N, STAGE)
tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout)
@@ -425,7 +425,7 @@ class GroupedGemmKernel:
ab_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# (EPI_TILE_M, EPI_TILE_N, STAGE)
sC: cute.struct.Align[
@@ -590,11 +590,11 @@ class GroupedGemmKernel:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -820,7 +820,7 @@ class SSDKernel:
num_threads=self.threads_per_cta,
)
tmem = utils.TmemAllocator(
smem_storage.tmem_holding_buf,
smem_storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=self.epilog_warp_id[0],
)
@@ -603,8 +603,6 @@ class MixedInputFusedMultiHeadAttentionDecode:
p_pipeline_ptr = smem.allocate_array(Int64, self.sp_stages * 2)
o_pipeline_ptr = smem.allocate_array(Int64, self.o_stages * 2)
assert smem._allocated_bytes <= self.mbarrier_reserved_bytes
# Declare named barriers
softmax_nbar_id = 1
mma_kq_nbar_id = 2
@@ -403,7 +403,7 @@ class MixedInputFusedMultiHeadAttentionPrefillD256:
s_corr_mbar_ptr: cute.struct.MemRange[Int64, self.qk_acc_stage * 2]
sum_mbar_ptr: cute.struct.MemRange[Int64, 2]
mma_o_mbar_ptr: cute.struct.MemRange[Int64, self.pv_acc_stage * 2]
tmem_dealloc_mbar_ptr: Int64
tmem_dealloc_mbar: Int64
tmem_holding_buf: Int32
self.shared_storage = SharedStorage
@@ -654,11 +654,11 @@ class MixedInputFusedMultiHeadAttentionPrefillD256:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=self.correction_warp_ids[0],
is_two_cta=True,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
pipeline_init_arrive(cluster_shape_mn=cluster_layout_vmnk, is_relaxed=True)
@@ -390,7 +390,7 @@ class MixedInputFusedMultiHeadAttentionPrefillD512:
p_mma_mbar_ptr: cute.struct.MemRange[Int64, self.qk_acc_stage * 2]
mma_o_mbar_ptr: cute.struct.MemRange[Int64, self.pv_acc_stage * 2]
swap_mbar_ptr: cute.struct.MemRange[Int64, self.swap_stage * 2]
tmem_dealloc_mbar_ptr: Int64
tmem_dealloc_mbar: Int64
tmem_holding_buf: Int32
self.shared_storage = SharedStorage
@@ -627,11 +627,11 @@ class MixedInputFusedMultiHeadAttentionPrefillD512:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=self.softmax_warp_ids[0],
is_two_cta=True,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
pipeline_init_arrive(cluster_shape_mn=cluster_layout_vmnk, is_relaxed=True)
@@ -624,7 +624,7 @@ class GroupedMixedInputGemmKernel:
tile_info_empty_mbar_ptr: cute.struct.MemRange[
cutlass.Int64, self.num_tile_info_stage
]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
self.shared_storage = SharedStorage
@@ -824,11 +824,11 @@ class GroupedMixedInputGemmKernel:
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_ptr_sync_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -518,7 +518,7 @@ class GroupedMixedInputGemmAccScaleKernel:
tile_info_empty_mbar_ptr: cute.struct.MemRange[
cutlass.Int64, self.num_tile_info_stage
]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
self.shared_storage = SharedStorage
@@ -708,11 +708,11 @@ class GroupedMixedInputGemmAccScaleKernel:
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_ptr_sync_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -618,7 +618,7 @@ class MixedInputGemmKernel:
]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# Tensor buffers
# (EPI_TILE_M, EPI_TILE_N, STAGE)
@@ -820,11 +820,11 @@ class MixedInputGemmKernel:
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_ptr_sync_barrier,
allocator_warp_id=self.epilog_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -568,7 +568,7 @@ class BlackwellMultiHeadLatentAttentionForwardFP16:
cutlass.Int64, self.load_pt_stage * 2
]
# Tmem dealloc cluster barrier
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
# Tmem holding buffer
tmem_holding_buf: cutlass.Int32
@@ -641,7 +641,6 @@ class BlackwellMultiHeadLatentAttentionForwardFP16:
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=self.cluster_shape_mnk,
smem=SplitKVKernelSharedStorage.size_in_bytes(),
stream=stream,
min_blocks_per_mp=1,
)
@@ -657,7 +656,6 @@ class BlackwellMultiHeadLatentAttentionForwardFP16:
).launch(
grid=(q_latent.shape[0], q_latent.shape[2], q_latent.shape[3]),
block=[self.threads_per_warp * self.num_compute_warps, 1, 1],
smem=MAX_SPLITS * self.acc_dtype.width // 8,
stream=stream,
min_blocks_per_mp=1,
)
@@ -838,11 +836,11 @@ class BlackwellMultiHeadLatentAttentionForwardFP16:
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_ptr_sync_bar,
allocator_warp_id=self.mma_warp_id,
is_two_cta=self.use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
load_q_pipeline = self.make_and_init_load_qkv_pipeline(
@@ -661,7 +661,7 @@ class BlackwellMultiHeadLatentAttentionForwardFP8:
]
# Tmem dealloc cluster barrier
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
# Tmem holding buffer
tmem_holding_buf: cutlass.Int32
@@ -707,7 +707,6 @@ class BlackwellMultiHeadLatentAttentionForwardFP8:
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=self.cluster_shape_mnk,
smem=SplitKVKernelSharedStorage.size_in_bytes(),
stream=stream,
min_blocks_per_mp=1,
)
@@ -723,7 +722,6 @@ class BlackwellMultiHeadLatentAttentionForwardFP8:
).launch(
grid=(q_latent.shape[0], q_latent.shape[2], q_latent.shape[3]),
block=[self.threads_per_warp * self.num_compute_warps, 1, 1],
smem=MAX_SPLITS * self.acc_dtype.width // 8,
stream=stream,
min_blocks_per_mp=1,
)
@@ -904,11 +902,11 @@ class BlackwellMultiHeadLatentAttentionForwardFP8:
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=self.tmem_ptr_sync_bar,
allocator_warp_id=self.mma_warp_id,
is_two_cta=self.use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
load_q_pipeline = self.make_and_init_load_qkv_pipeline(
@@ -43,9 +43,9 @@ from cutlass.cute.runtime import from_dlpack
from dataclasses import dataclass, field
"""
This example provides an experimental implementation of the SM103 batched 3xFP4 blockscaled GEMM kernel, please note that the APIs and implementation details related to this kernel may change in future releases.
This example provides an experimental implementation of the SM103 batched FP4 Ultra blockscaled GEMM kernel, please note that the APIs and implementation details related to this kernel may change in future releases.
A high-performance persistent batched 3xFP4 blockscaled GEMM example for the NVIDIA Blackwell SM103 architecture
A high-performance persistent batched FP4 Ultra blockscaled GEMM example for the NVIDIA Blackwell SM103 architecture
using CUTE DSL.
- Matrix A is MxKxL, L is batch dimension, A can only be row-major("K") for MXF4/NVF4 input type
- Matrix B is NxKxL, L is batch dimension, B can only be row-major("K") for MXF4/NVF4 input type
@@ -166,7 +166,7 @@ class Sm103BlockScaledPersistentDenseGemmKernel:
cluster_shape_mn: Tuple[int, int],
use_tma_store: bool,
):
"""Initializes the configuration for a Blackwell SM103 3xFP4 GEMM kernel.
"""Initializes the configuration for a Blackwell SM103 FP4 Ultra GEMM kernel.
This configuration includes several key aspects:
@@ -603,7 +603,7 @@ class Sm103BlockScaledPersistentDenseGemmKernel:
sf_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_sf_stage]
acc_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
acc_empty_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
# (MMA, MMA_M, MMA_K, STAGE)
sA: cute.struct.Align[
@@ -800,11 +800,11 @@ class Sm103BlockScaledPersistentDenseGemmKernel:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=self.epilogue_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -1810,7 +1810,7 @@ class Sm103BlockScaledPersistentDenseGemmKernel:
mma_tiler_mn: Tuple[int, int],
a_source: tcgen05.OperandSource = tcgen05.OperandSource.SMEM,
) -> cute.TiledMma:
"""Create a blockscaled trivial tiled MMA for SM103 (3xFP4), K fixed to 96.
"""Create a blockscaled trivial tiled MMA for SM103 (FP4 Ultra), K fixed to 96.
Returns a tcgen05 MMA configured for the given (M, N) tiler and CTA group.
@@ -2653,7 +2653,7 @@ def run(
:return: Execution time of the GEMM kernel
:rtype: float
"""
print(f"Running Sm103 Persistent 3xfp4 Dense BlockScaled GEMM test with:")
print(f"Running Sm103 Persistent FP4 Ultra Dense BlockScaled GEMM test with:")
print(f"mnkl: {mnkl}")
print(f"AB dtype: {ab_dtype}, SF dtype: {sf_dtype}, SF Vec size: {sf_vec_size}")
print(f"C dtype: {c_dtype}")
@@ -2954,7 +2954,7 @@ if __name__ == "__main__":
)
parser = argparse.ArgumentParser(
description="Example of Sm103 3xfp4 Dense Persistent BlockScaled GEMM."
description="Example of Sm103 FP4 Ultra Dense Persistent BlockScaled GEMM."
)
parser.add_argument(
@@ -8,13 +8,10 @@
# without an express license agreement from NVIDIA CORPORATION or
# its affiliates is strictly prohibited.
import argparse
from typing import Tuple, Type, Callable
from functools import partial, lru_cache
from typing import Tuple
import cutlass
from cutlass import Numeric
import cutlass.cute as cute
import cutlass.utils as utils
import cutlass.pipeline as pipeline
@@ -69,7 +66,6 @@ def kernel(
a_smem_layout: cute.ComposedLayout,
b_smem_layout: cute.ComposedLayout,
):
# Current thread/warp/block coordinates
tidx, _, _ = cute.arch.thread_idx()
warp_idx = cute.arch.warp_idx()
@@ -103,7 +99,7 @@ def kernel(
num_threads=threads_per_cta,
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
)
num_tmem_cols = 512
@@ -143,15 +139,15 @@ def kernel(
# (bM, bN)
gC = cute.local_tile(mC_mnl, mma_tiler_mnk, mma_coord_mnk, proj=(1, 1, None))
thr_mma = tiled_mma.get_slice(0)
# (MMA, MMA_M, MMA_K, RestK)
# (MMA, MMA_M, MMA_K)
tCgA = thr_mma.partition_A(gA)
# (MMA, MMA_N, MMA_K, RestK)
# (MMA, MMA_N, MMA_K)
tCgB = thr_mma.partition_B(gB)
# (MMA, MMA_M, MMA_N)
tCgC = thr_mma.partition_C(gC)
# (MMA, MMA_M, MMA_K, STAGE)
# (MMA, MMA_M, MMA_K)
tCrA = tiled_mma.make_fragment_A(sA)
# (MMA, MMA_N, MMA_K, STAGE)
# (MMA, MMA_N, MMA_K)
tCrB = tiled_mma.make_fragment_B(sB)
# (MMA, MMA_M, MMA_N)
acc_shape = tiled_mma.partition_shape_C(mma_tiler_mnk[:2])
@@ -199,14 +195,14 @@ def kernel(
tmem_thr_copy = tmem_tiled_copy.get_slice(tidx)
# (TmemCpy,NumTmemCpy,NumTiles)
tCtC = tmem_thr_copy.partition_S(tCtAcc_epi)
tDtC = tmem_thr_copy.partition_S(tCtAcc_epi)
# (TmemCpy,NumTmemCpy,NumTiles)
tCgC = tmem_thr_copy.partition_D(gC_epi)
tDgC = tmem_thr_copy.partition_D(gC_epi)
# (TmemCpy,NumTmemCpy)
tCrAcc = cute.make_rmem_tensor(tCgC[None, None, 0].shape, acc_dtype)
tCrAcc = cute.make_rmem_tensor(tDgC[None, None, 0].shape, acc_dtype)
# (TmemCpy,NumTmemCpy)
tCrC = cute.make_rmem_tensor(tCgC[None, None, 0].shape, io_dtype)
tCrC = cute.make_rmem_tensor(tDgC[None, None, 0].shape, io_dtype)
#
# 2. Main loop
@@ -233,8 +229,6 @@ def kernel(
# Execute one K-block worth of MMA instructions
ab_full = ab_consumer.wait_and_advance()
# tCtAcc += tCrA * tCrB
num_k_blocks = cute.size(tCrA, mode=[2])
for k_block_idx in cutlass.range_constexpr(num_k_blocks):
k_block_coord = (None, None, k_block_idx, ab_full.index)
@@ -265,10 +259,10 @@ def kernel(
# TMEM -> RMEM -> GEMM
# Sub-tiling for better instruction-level parallelism
for i in cutlass.range(cute.size(tCtC, mode=[2])):
cute.copy(tmem_tiled_copy, tCtC[None, None, i], tCrAcc)
for i in cutlass.range(cute.size(tDtC, mode=[2])):
cute.copy(tmem_tiled_copy, tDtC[None, None, i], tCrAcc)
tCrC.store(tCrAcc.load().to(io_dtype))
cute.autovec_copy(tCrC, tCgC[None, None, i])
cute.autovec_copy(tCrC, tDgC[None, None, i])
acc_full.release()
# Deallocate TMEM
@@ -350,44 +344,10 @@ def host_function(a: cute.Tensor, b: cute.Tensor, c: cute.Tensor):
)
@lru_cache(maxsize=1)
def prepare_run(
callable: Callable,
m: int,
n: int,
k: int,
a_dtype: Type[Numeric],
b_dtype: Type[Numeric],
c_dtype: Type[Numeric],
) -> tuple[Callable, tuple]:
import cutlass.torch as cutlass_torch
a, b, c = cutlass_torch.prepare_tensors_for_gemm(
(m, n, k), a_dtype, b_dtype, c_dtype
)
a_ = (
from_dlpack(a, assumed_align=32)
.mark_layout_dynamic(leading_dim=1)
.mark_compact_shape_dynamic(mode=1, divisibility=k)
)
b_ = (
from_dlpack(b, assumed_align=32)
.mark_layout_dynamic(leading_dim=1)
.mark_compact_shape_dynamic(mode=1, divisibility=k)
)
c_ = (
from_dlpack(c, assumed_align=32)
.mark_layout_dynamic(leading_dim=1)
.mark_compact_shape_dynamic(mode=1, divisibility=n)
)
compiled_fn = cute.compile(callable, a_, b_, c_, options="--generate-line-info")
return partial(compiled_fn, a_, b_, c_), (a, b, c)
def run_dense_gemm(
mnk: Tuple[int, int, int],
tolerance: float,
) -> None:
):
global torch, cutlass_torch
import torch
import cutlass.torch as cutlass_torch
@@ -402,23 +362,48 @@ def run_dense_gemm(
m, n, k = mnk
torch.manual_seed(1111)
run_fn, (a, b, c) = prepare_run(
host_function, m, n, k, io_dtype, io_dtype, io_dtype
# Make K-major tensors (torch tensors are row-major)
def make_tensors(mn, k, dtype):
shape = (mn, k)
return (
torch.empty(*shape, dtype=torch.int32)
.random_(-2, 2)
.to(dtype=dtype, device="cuda")
)
a = make_tensors(m, k, cutlass_torch.dtype(io_dtype))
b = make_tensors(n, k, cutlass_torch.dtype(io_dtype))
c = make_tensors(m, n, cutlass_torch.dtype(io_dtype))
a_tensor = (
from_dlpack(a, assumed_align=32)
.mark_layout_dynamic(leading_dim=1)
.mark_compact_shape_dynamic(mode=1, divisibility=k)
)
b_tensor = (
from_dlpack(b, assumed_align=32)
.mark_layout_dynamic(leading_dim=1)
.mark_compact_shape_dynamic(mode=1, divisibility=k)
)
c_tensor = (
from_dlpack(c, assumed_align=32)
.mark_layout_dynamic(leading_dim=1)
.mark_compact_shape_dynamic(mode=1, divisibility=n)
)
# Entry point to the host JIT function
run_fn()
host_function(a_tensor, b_tensor, c_tensor, no_cache=True)
# Compute reference result and verify
ref = torch.einsum("mk,nk->mn", a.to(torch.float32), b.to(torch.float32))
ref = (torch.einsum("mk,nk->mn", a.to(torch.float32), b.to(torch.float32))).cpu()
torch.testing.assert_close(
c, ref.to(cutlass_torch.dtype(io_dtype)), atol=tolerance, rtol=1e-05
c.cpu(), ref.to(cutlass_torch.dtype(io_dtype)), atol=tolerance, rtol=1e-05
)
if __name__ == "__main__":
def parse_comma_separated_ints(s: str) -> list[int]:
def parse_comma_separated_ints(s: str):
try:
return [int(x.strip()) for x in s.split(",")]
except ValueError:
@@ -443,14 +428,14 @@ if __name__ == "__main__":
parser.add_argument(
"--tolerance", type=float, default=1e-01, help="Tolerance for validation"
)
args = parser.parse_args()
if len(args.mnk) != 3:
parser.error("--mnk must contain exactly 3 values")
if args.mnk[0] % mma_tiler_mnk[0] != 0 or args.mnk[1] % mma_tiler_mnk[1] != 0:
parser.error("m n must be divisible by mma_tiler_mn")
run_dense_gemm(args.mnk, args.tolerance)
run_dense_gemm(
args.mnk,
args.tolerance,
)
print("PASS")
@@ -65,8 +65,7 @@ Constraints for this example:
io_dtype = cutlass.Float16
acc_dtype = cutlass.Float32
use_2cta_instrs = True
cluster_shape_mnk = (2, 1, 1) if use_2cta_instrs else (1, 1, 1)
cluster_shape_mnk = (2, 1, 1)
mma_inst_shape_mnk = (256, 256, 16)
mma_tiler_mnk = (256, 256, 64)
threads_per_cta = 128
@@ -96,7 +95,6 @@ def kernel(
b_smem_layout: cute.ComposedLayout,
cta_layout_vmnk: cute.Layout,
):
# Current thread/warp/block coordinates
tidx, _, _ = cute.arch.thread_idx()
warp_idx = cute.arch.warp_idx()
@@ -174,15 +172,15 @@ def kernel(
# (bM, bN)
gC = cute.local_tile(mC_mnl, mma_tiler_mnk, mma_coord_mnk, proj=(1, 1, None))
thr_mma = tiled_mma.get_slice(mma_coord_vmnk[0])
# (MMA, MMA_M, MMA_K, RestK)
# (MMA, MMA_M, MMA_K)
tCgA = thr_mma.partition_A(gA)
# (MMA, MMA_N, MMA_K, RestK)
# (MMA, MMA_N, MMA_K)
tCgB = thr_mma.partition_B(gB)
# (MMA, MMA_M, MMA_N)
tCgC = thr_mma.partition_C(gC)
# (MMA, MMA_M, MMA_K, STAGE)
# (MMA, MMA_M, MMA_K)
tCrA = tiled_mma.make_fragment_A(sA)
# (MMA, MMA_N, MMA_K, STAGE)
# (MMA, MMA_N, MMA_K)
tCrB = tiled_mma.make_fragment_B(sB)
# (MMA, MMA_M, MMA_N)
acc_shape = tiled_mma.partition_shape_C(mma_tiler_mnk[:2])
@@ -217,10 +215,10 @@ def kernel(
num_threads=threads_per_cta,
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
is_two_cta=cute.size(cta_layout_vmnk, mode=[0]) > 1,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
num_tmem_cols = 512
tmem.allocate(num_tmem_cols)
@@ -232,7 +230,7 @@ def kernel(
# Swap the pointer in tCtAcc
tCtAcc = cute.make_tensor(tmem_ptr, tCtAcc.layout)
subtile_cnt = 1 if mma_tiler_mnk[0] == 64 else 4
subtile_cnt = 4
# (EpiTile)
epi_tiler = (
(cute.size(tCtAcc, mode=[0, 0]), cute.size(tCtAcc, mode=[0, 1]) // subtile_cnt),
@@ -244,24 +242,21 @@ def kernel(
# Every thread loads 64 x fp32
tmem_atom = cute.make_copy_atom(
tcgen05.Ld16x256bOp(tcgen05.Repetition.x8)
if mma_tiler_mnk[0] == 64
else tcgen05.Ld32x32bOp(tcgen05.Repetition.x64),
tcgen05.Ld32x32bOp(tcgen05.Repetition.x64),
cutlass.Float32,
)
tmem_tiled_copy = tcgen05.make_tmem_copy(tmem_atom, tCtAcc_epi[None, 0])
tmem_thr_copy = tmem_tiled_copy.get_slice(tidx)
# (TmemCpy,NumTmemCpy,NumTiles)
tCtC = tmem_thr_copy.partition_S(tCtAcc_epi)
tDtC = tmem_thr_copy.partition_S(tCtAcc_epi)
# (TmemCpy,NumTmemCpy,NumTiles)
tCgC = tmem_thr_copy.partition_D(gC_epi)
tDgC = tmem_thr_copy.partition_D(gC_epi)
# (TmemCpy,NumTmemCpy)
tCrAcc = cute.make_rmem_tensor(tCgC[None, None, 0].shape, acc_dtype)
tCrAcc = cute.make_rmem_tensor(tDgC[None, None, 0].shape, acc_dtype)
# (TmemCpy,NumTmemCpy)
tCrC = cute.make_rmem_tensor(tCgC[None, None, 0].shape, io_dtype)
tCrC = cute.make_rmem_tensor(tDgC[None, None, 0].shape, io_dtype)
#
# 2. Main loop
@@ -271,8 +266,8 @@ def kernel(
if warp_idx == 0:
# Wait for a empty accumulator buffer
if is_leader_cta:
acc_producer.acquire()
for k_tile in cutlass.range(num_k_tiles, prefetch_stages=ab_stages - 2):
acc_producer.acquire_and_advance()
for _ in cutlass.range(num_k_tiles, prefetch_stages=ab_stages - 2):
# Issue TMA loads
ab_empty = ab_producer.acquire_and_advance()
cute.copy(
@@ -310,7 +305,6 @@ def kernel(
# Signal that the accumulator is fully computed
if is_leader_cta:
acc_producer.commit()
acc_producer.advance()
#
# 3. Epilogue
@@ -321,13 +315,12 @@ def kernel(
# Wait for the accumulator buffer to be full
acc_full = acc_consumer.wait_and_advance()
# TMEM -> RMEM -> GEMM
# Sub-tiling for better instruction-level parallelism
for i in cutlass.range(cute.size(tCtC, mode=[2])):
cute.copy(tmem_tiled_copy, tCtC[None, None, i], tCrAcc)
for i in cutlass.range(cute.size(tDtC, mode=[2])):
cute.copy(tmem_tiled_copy, tDtC[None, None, i], tCrAcc)
tCrC.store(tCrAcc.load().to(io_dtype))
cute.autovec_copy(tCrC, tCgC[None, None, i])
cute.autovec_copy(tCrC, tDgC[None, None, i])
acc_full.release()
# Ensure used buffers are properly synchronized before producer exit.
@@ -353,7 +346,7 @@ def host_function(
io_dtype,
acc_dtype,
mma_inst_shape_mnk,
tcgen05.CtaGroup.TWO if use_2cta_instrs else tcgen05.CtaGroup.ONE,
tcgen05.CtaGroup.TWO,
tcgen05.OperandSource.SMEM,
tcgen05.OperandMajorMode.K,
tcgen05.OperandMajorMode.K,
@@ -381,16 +374,14 @@ def host_function(
cta_layout_vmnk = cute.tiled_divide(cta_layout_mnk, (tiled_mma.thr_id,))
# Construct TMA load atoms
op = cute.nvgpu.cpasync.CopyBulkTensorTileG2SMulticastOp(
tcgen05.CtaGroup.TWO if use_2cta_instrs else tcgen05.CtaGroup.ONE
)
op = cute.nvgpu.cpasync.CopyBulkTensorTileG2SMulticastOp(tcgen05.CtaGroup.TWO)
a_tma_atom, a_tma_tensor = cute.nvgpu.make_tiled_tma_atom_A(
op,
a,
a_smem_layout_one_stage,
mma_tiler_mnk,
tiled_mma,
cta_layout_vmnk.shape,
cta_layout_vmnk.shape, # take the layout and extract the shape internally
)
b_tma_atom, b_tma_tensor = cute.nvgpu.make_tiled_tma_atom_B(
op,
@@ -403,8 +394,7 @@ def host_function(
grid_shape = cute.round_up(
cute.ceil_div(
(*c.layout.shape, 1),
(mma_tiler_mnk[0] // (2 if use_2cta_instrs else 1), *mma_tiler_mnk[1:]),
(*c.layout.shape, 1), (mma_tiler_mnk[0] // 2, *mma_tiler_mnk[1:])
),
cluster_shape_mnk,
)
@@ -981,7 +981,7 @@ def run_dense_gemm(
import cutlass.torch as cutlass_torch
print("===================================================================")
print("Running Blackwell fp16 GEMM example 4 (with MIX CGA support):")
print("Running Blackwell fp16 GEMM example 4 (with MIX cluster size support):")
print(f" mnk: {mnk}")
print(f" tolerance: {tolerance}")
print(f" Preferred cluster shape: {preferred_cluster_shape_mnk}")
@@ -500,7 +500,7 @@ class Sm100BlockScaledDenseGemmKernel:
num_threads=self.threads_per_cta,
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
)
tmem.allocate(self.num_tmem_alloc_cols)
@@ -436,7 +436,7 @@ class Sm100BlockScaledDenseGemmKernel:
class SharedStorage:
ab_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage * 2]
acc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_acc_stage * 2]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
smem = utils.SmemAllocator()
@@ -638,10 +638,10 @@ class Sm100BlockScaledDenseGemmKernel:
num_threads=self.threads_per_cta,
)
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
is_two_cta=cute.size(cta_layout_vmnk, mode=[0]) > 1,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
tmem.allocate(self.num_tmem_alloc_cols)
tmem.wait_for_alloc()
@@ -581,7 +581,7 @@ class PersistentDenseGemmKernel:
acc_full_mbar_ptr: cute.struct.MemRange[
cutlass.Int64, self.num_acc_stage * 2
]
tmem_dealloc_mbar_ptr: cutlass.Int64
tmem_dealloc_mbar: cutlass.Int64
tmem_holding_buf: cutlass.Int32
smem = utils.SmemAllocator()
@@ -632,11 +632,11 @@ class PersistentDenseGemmKernel:
)
# Tensor memory dealloc barrier init
tmem = utils.TmemAllocator(
storage.tmem_holding_buf,
storage.tmem_holding_buf.ptr,
barrier_for_retrieve=tmem_alloc_barrier,
allocator_warp_id=self.epilogue_warp_id[0],
is_two_cta=use_2cta_instrs,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr,
two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr,
)
# Cluster arrive after barrier init
@@ -3,7 +3,7 @@
#
# Use of this software is governed by the terms and conditions of the
# NVIDIA End User License Agreement (EULA), available at:
# https://docs.nvidia.com/cutlass/media/docs/pythonDSL/license.html
# https://docs.nvidia.com/cutlass/latest/media/docs/pythonDSL/license.html
#
# Any use, reproduction, disclosure, or distribution of this software
# and related documentation outside the scope permitted by the EULA
-1
View File
@@ -527,7 +527,6 @@ class HopperFusedMultiHeadAttentionForward:
grid=grid,
block=[self.threads_per_cta, 1, 1],
cluster=self.cluster_shape_mnk,
smem=self.shared_storage.size_in_bytes(),
stream=stream,
min_blocks_per_mp=1,
)
@@ -107,7 +107,7 @@ Constraints (same as dense_gemm_persistent.py plus):
* Cluster shape M/N: power of 2, total <= 4
* Contiguous dim must be 16-byte aligned
Debug environment knobs:
Debug environment options:
* `GROUPED_GEMM_FORCE_CUTE_COPY=1`
Disable the non-mcast NVVM TMA load path and always use `cute.copy`.
"""
@@ -30,7 +30,6 @@ from functools import partial
import jax
import jax.numpy as jnp
import cutlass
import cutlass.cute as cute
import cutlass.jax as cjax
import cuda.bindings.driver as cuda
@@ -140,12 +139,12 @@ if __name__ == "__main__":
def run_cutlass_kernel(a, b, x, y):
call = cjax.cutlass_call(
launch_jax_wrapper,
# Jax requires output shapes/dtype information for each output
# Describe the shape and dtype of each output buffer.
output_shape_dtype=(
jax.ShapeDtypeStruct(a.shape, a.dtype),
jax.ShapeDtypeStruct(b.shape, a.dtype),
),
# Static jit arguments are passed via additional keyword arguments
# Static jit arguments are passed via additional keyword arguments.
x=x,
y=y,
)
@@ -165,12 +164,11 @@ if __name__ == "__main__":
# to the kernel. Alternatively you can wrap using another separate cute.jit
# function.
lambda stream, a, b, c, d, *, x, y: launch(a, b, x, y, c, d, stream),
# Jax requires output shapes/dtype information for each output
output_shape_dtype=(
jax.ShapeDtypeStruct(a.shape, a.dtype),
jax.ShapeDtypeStruct(b.shape, a.dtype),
),
# Static jit arguments are passed via additional keyword arguments
# Static jit arguments are passed via additional keyword arguments.
x=x,
y=y,
)
@@ -191,11 +189,12 @@ if __name__ == "__main__":
jax.ShapeDtypeStruct(a.shape, a.dtype),
jax.ShapeDtypeStruct(b.shape, a.dtype),
),
# By default cutlass_call will treat all tensors as dynamic shape.
# By default cutlass_call treats all tensors as dynamic shape.
# Dynamic shapes are often expected for kernels so this default ensures
# the broadest support. If you know that a kernel can accept fully static
# tensors then you can enable this flag to pass all tensors shapes and
# layouts known at compile time.
# tensors then you can enable this flag to compile all tensor shapes and
# layouts as constexpr values known at compile time.
# Individual tensors may opt out via .mark_layout_dynamic().
use_static_tensors=True,
x=x,
y=y,
@@ -209,19 +208,15 @@ if __name__ == "__main__":
@partial(jax.jit, static_argnums=[2, 3])
def run_cutlass_kernel_with_modes(a, b, x, y):
# input_spec and output_spec accept TensorSpec values to attach layout
# metadata to tensors. mode remaps the logical dimension order seen by
# the kernel. static=True compiles that tensor's layout as constexpr.
call = cjax.cutlass_call(
lambda stream, a, b, c, d, *, x, y: launch(a, b, x, y, c, d, stream),
output_shape_dtype=(
jax.ShapeDtypeStruct(a.shape, a.dtype),
jax.ShapeDtypeStruct(b.shape, a.dtype),
),
# The modes of the layout for each tensor can be specified using the
# TensorSpec. By default modes will align with the physical layout
# but can be mapped to specific index position. If None is passed
# then the default mode is assumed for that tensor.
#
# Individual static/dynamic settings may also be applied. For example
# a specific tensor can be marked to have static shape.
input_spec=(
cjax.TensorSpec(mode=(1, 0, 2), static=True),
cjax.TensorSpec(mode=(3, 1, 2, 0)),
@@ -245,9 +240,8 @@ if __name__ == "__main__":
jax.ShapeDtypeStruct(a.shape, a.dtype),
jax.ShapeDtypeStruct(b.shape, b.dtype),
),
# Can specify the input tensors that are aliasing outputs of this call.
# To avoid allocating separate output buffers. This is useful for kernels
# that update a tensor.
# Map input indices to output indices so XLA can reuse the input
# buffers for the outputs, avoiding extra allocations.
input_output_aliases={0: 0, 1: 1},
x=x,
y=y,
@@ -26,45 +26,45 @@
# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
import pytest
from functools import partial
import argparse
import cuda.bindings.driver as cuda
import cutlass
import cutlass.cute as cute
import jax
import jax.numpy as jnp
from jax import export
from cutlass.jax import cutlass_call, get_export_disabled_safety_checks
from cutlass.jax.testing import create_tensor
"""
Examples of using jax.export APIs with functions using cutlass_call.
This example demonstrates the use of jax.export with CuTe DSL kernel. It assumes
familiarity with CuTe DSL concepts such as layouts and dynamic shapes as well as
Jax's exporting and serialization features:
This example demonstrates three export modes:
1. Concrete shapes -- shapes are fixed constants baked into the export.
2. Unconstrained symbolic shapes ("a, b")
3. Constrained symbolic shapes ("32*M, 16*N")
The JAX function being exported is the same in all three cases; only the
shape specification passed to jax.export differs.
It assumes familiarity with CuTe DSL concepts such as layouts and dynamic shapes
as well as JAX's exporting and serialization features:
https://docs.jax.dev/en/latest/export/index.html#export
To run this example:
.. code-block:: bash
# Run with defaults
python examples/jax/cutlass_call_export.py
python examples/jax/cutlass_call_export.py --M 512 --N 256
# Run with shape (1024, 512)
python examples/jax/cutlass_call_export.py --M 1024 --N 512
# Export with symbolic shapes.
python examples/jax/cutlass_call_export.py --export_symbolic
"""
import argparse
import cuda.bindings.driver as cuda
import cutlass.cute as cute
import jax
import jax.numpy as jnp
from jax import export
from cutlass.jax import cutlass_call, get_export_disabled_safety_checks, TensorSpec
from cutlass.jax.testing import create_tensor
# Simple element-wise addition kernel: gC[i,j] = gA[i,j] + gB[i,j]
@cute.kernel
def kernel(gA: cute.Tensor, gB: cute.Tensor, gC: cute.Tensor):
tidx, _, _ = cute.arch.thread_idx()
@@ -84,9 +84,6 @@ def kernel(gA: cute.Tensor, gB: cute.Tensor, gC: cute.Tensor):
@cute.jit
def launch(stream: cuda.CUstream, mA: cute.Tensor, mB: cute.Tensor, mC: cute.Tensor):
print("mA: ", mA.layout)
print("mB: ", mB.layout)
print("mC: ", mC.layout)
num_threads_per_block = 256
m, n = mA.shape
kernel(mA, mB, mC).launch(
@@ -96,63 +93,100 @@ def launch(stream: cuda.CUstream, mA: cute.Tensor, mB: cute.Tensor, mC: cute.Ten
)
def run_example(M, N, export_symbolic_shapes):
def _export_and_run(f, ref_f, input_shape_dtype, run_shapes):
"""Export f, serialize/deserialize, then run on each shape in run_shapes.
Both inputs (a, b) are assumed to share the same input_shape_dtype.
"""
print(f"Exporting with input signature: ({input_shape_dtype}, {input_shape_dtype})")
# jax.export can be used to export a jit function containing cutlass_call.
# CUTLASS custom call targets are not on JAX's built-in stable custom-call
# allowlist, so we pass them via disabled_checks to suppress that safety check.
exported = jax.export.export(f, disabled_checks=get_export_disabled_safety_checks())
traced = exported(input_shape_dtype, input_shape_dtype)
blob = traced.serialize()
print(f"Serialized computation is {len(blob)} bytes.")
rehydrated = export.deserialize(blob)
key = jax.random.key(1123)
a_key, b_key = jax.random.split(key, 2)
for shape in run_shapes:
a = create_tensor(shape, dtype=jnp.float32, key=a_key)
b = create_tensor(shape, dtype=jnp.float32, key=b_key)
c = rehydrated.call(a, b)
assert jnp.allclose(c, ref_f(a, b)), f"Mismatch at shape {shape}"
print(f" shape {shape}: OK")
def run_example(M, N):
@jax.jit
def ref_f(a, b):
return jax.nn.sigmoid(a + b)
# The same JAX function is used in all three examples below. The export
# mode is determined entirely by the shape spec passed to jax.export.
@jax.jit
def f(a, b):
call = cutlass_call(launch, output_shape_dtype=a)
return jax.nn.sigmoid(call(a, b))
# ── 1. Concrete shapes ────────────────────────────────────────────────────
# Shapes are fixed constants baked into the export. The deserialized
# computation only accepts exactly these dimensions at runtime.
print("\nConcrete shapes:")
input_shape_dtype = jax.ShapeDtypeStruct((M, N), jnp.float32)
_export_and_run(
f,
ref_f,
input_shape_dtype,
run_shapes=[(M, N)], # concrete exports reject any other shape
)
# ── 2. Unconstrained symbolic shapes ─────────────────────────────────────
# Both dimensions are fully dynamic. The exported computation accepts any
# (M, N) at runtime without recompilation.
print("\nUnconstrained symbolic shapes:")
a_sym, b_sym = export.symbolic_shape("a, b")
input_shape_dtype = jax.ShapeDtypeStruct((a_sym, b_sym), jnp.float32)
_export_and_run(
f,
ref_f,
input_shape_dtype,
run_shapes=[(M, N), (M * 2, N * 4), (M * 4, N * 4)],
)
# ── 3. Constrained symbolic shapes (divisibility) ─────────────────────────
# Shapes are declared as multiples of a tile size via TensorSpec.divisibility.
# The symbolic expression "32*M, 16*N" tells jax.export that dim 0 is always
# a multiple of 32 and dim 1 is always a multiple of 16. This lets the
# compiler generate more efficient code (e.g. no remainder handling).
# Runtime shapes must satisfy these divisibility constraints.
print("\nConstrained symbolic shapes:")
@jax.jit
def ref_f(a, b):
return jax.nn.sigmoid(a + b)
def f_divisible(a, b):
spec = TensorSpec(divisibility=(32, 16))
call = cutlass_call(
launch,
output_shape_dtype=a,
input_spec=(spec, spec),
output_spec=spec,
)
return jax.nn.sigmoid(call(a, b))
# Symbolic or partially shapes are supported by cutlass_call and cute.Tensor
# This allows export of functions calling Cut eDSL kernels w/o having to re-compile
# the kernel for each new shape.
if export_symbolic_shapes:
a, b = export.symbolic_shape("a, b")
export_shape_dtype = jax.ShapeDtypeStruct((a, b), jnp.float32)
else:
export_shape_dtype = jax.ShapeDtypeStruct((M, N), jnp.float32)
print("Exporting with input signature: ")
print(f"({export_shape_dtype}, {export_shape_dtype})")
# jax.export can be used to export a jit function containing cutlass_call.
# The function get_export_disabled_safety_checks() returns a list of custom
# call targets that are used by cutlass_call not part of Jax's built-in
# list of stable custom calls.
exported = jax.export.export(f, disabled_checks=get_export_disabled_safety_checks())
traced = exported(export_shape_dtype, export_shape_dtype)
# Serialize the computation to a byte blob.
blob = traced.serialize()
print(f"Serialized computation is {len(blob)} bytes.")
# Deserialize and run
rehydrated = export.deserialize(blob)
key = jax.random.key(1123)
a_key, b_key = jax.random.split(key, 2)
a = create_tensor((M, N), dtype=jnp.float32, key=a_key)
b = create_tensor((M, N), dtype=jnp.float32, key=b_key)
c = rehydrated.call(a, b)
assert jnp.allclose(c, ref_f(a, b))
# If the computation was exported with dynamic shapes then we can also
# call it with different shapes. The kernel will not be re-compiled
# even though the shapes are changing.
if export_symbolic_shapes:
a = create_tensor((M * 2, N * 4), dtype=jnp.float32, key=a_key)
b = create_tensor((M * 2, N * 4), dtype=jnp.float32, key=b_key)
c = rehydrated.call(a, b)
assert jnp.allclose(c, ref_f(a, b))
a = create_tensor((M * 4, N * 4), dtype=jnp.float32, key=a_key)
b = create_tensor((M * 4, N * 4), dtype=jnp.float32, key=b_key)
c = rehydrated.call(a, b)
assert jnp.allclose(c, ref_f(a, b))
m_sym, n_sym = export.symbolic_shape("32*M, 16*N")
input_shape_dtype = jax.ShapeDtypeStruct((m_sym, n_sym), jnp.float32)
_export_and_run(
f_divisible,
ref_f,
input_shape_dtype,
run_shapes=[(M, N), (M * 2, N * 2), (M * 4, N * 4)],
)
if __name__ == "__main__":
@@ -161,8 +195,7 @@ if __name__ == "__main__":
)
parser.add_argument("--M", default=512, type=int)
parser.add_argument("--N", default=256, type=int)
parser.add_argument("--export_symbolic", action="store_true")
args = parser.parse_args()
run_example(args.M, args.N, args.export_symbolic)
run_example(args.M, args.N)
print("PASS")
@@ -27,14 +27,12 @@
# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
from functools import partial
import argparse
import jax
import jax.numpy as jnp
from jax.sharding import NamedSharding, PartitionSpec as P, AxisType
from jax.experimental.custom_partitioning import custom_partitioning
import cutlass
import cutlass.cute as cute
import cutlass.jax as cjax
from cutlass.jax.testing import create_tensor
@@ -30,7 +30,7 @@
import argparse
import operator
from functools import partial
from typing import List, Type
from typing import List
import cuda.bindings.driver as cuda
import cutlass
@@ -78,7 +78,7 @@
},
{
"cell_type": "code",
"execution_count": 2,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
@@ -114,7 +114,7 @@
" ]\n",
"\n",
" synced_producer_consumer(SharedStorage, res).launch(\n",
" grid=(1, 1, 1), block=(64, 1, 1), smem=SharedStorage.size_in_bytes()\n",
" grid=(1, 1, 1), block=(64, 1, 1)\n",
" )\n",
"\n",
"\n",
@@ -455,7 +455,7 @@
},
{
"cell_type": "code",
"execution_count": 6,
"execution_count": null,
"metadata": {},
"outputs": [],
"source": [
@@ -527,7 +527,7 @@
" ]\n",
"\n",
" async_pipeline_staged_kernel(SharedStorage, res, staging).launch(\n",
" grid=(1, 1, 1), block=(64, 1, 1), smem=SharedStorage.size_in_bytes()\n",
" grid=(1, 1, 1), block=(64, 1, 1)\n",
" )\n",
"\n",
"\n",