CUTLASS 3.8 Release (#2059)

* CUTLASS 3.8 Release

* update

* Update README.md

* Revert "Update README.md"

This reverts commit b353e36fe83e0815f99b44e46c0c95494c44726b.

* update

* update

---------

Co-authored-by: Haicheng Wu <57973641+hwu36@users.noreply.github.com>
Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
This commit is contained in:
mihir-awatramani
2025-01-25 02:44:06 -05:00
committed by GitHub
co-authored by Haicheng Wu Haicheng Wu
parent 9eb01fa0b0
commit 389e493055
290 changed files with 91222 additions and 291 deletions
@@ -0,0 +1,193 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2025 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.
*
**************************************************************************************************/
//
//
#pragma once
#include "cutlass/layout/tensor.h"
#include "cute/atom/copy_traits_sm100_im2col.hpp"
#include "cutlass/arch/mma.h"
#include "cutlass/conv/convolution.h"
#include "cutlass/conv/dispatch_policy.hpp"
#include "cutlass/detail/layout.hpp"
#include "cutlass/conv/collective/builders/sm90_common.inl"
#include "cutlass/gemm/collective/builders/sm100_common.inl"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass::conv::collective::detail {
/////////////////////////////////////////////////////////////////////////////////////////////////
// Collective tile traits struct that serves as a type list containing a tensor's mem layouts and atoms
template<
class GmemTiledCopy_,
class SmemLayoutAtom_,
class TmemLayoutAtom_ = void
>
struct Sm100ImplicitGemmTileTraits {
using GmemTiledCopy = GmemTiledCopy_;
using SmemLayoutAtom = SmemLayoutAtom_;
using TmemLayoutAtom = TmemLayoutAtom_;
};
template <class ClusterShapeMNK, class AtomThrId>
constexpr auto
sm100_cluster_shape_to_im2col_tma_atom_A(ClusterShapeMNK cluster_shape_mnk, AtomThrId atom_thr_id) {
static_assert(cute::rank(cluster_shape_mnk) == 3);
constexpr bool IsDynamicCluster = not cute::is_static_v<ClusterShapeMNK>;
if constexpr (cute::size(atom_thr_id) == 2) {
if constexpr (!IsDynamicCluster) {
static_assert(cute::size<0>(cluster_shape_mnk) % 2 == 0, "Cluster shape not divisible by MMA size");
if constexpr (cute::size<1>(cluster_shape_mnk) == 1) {
return cute::SM100_TMA_2SM_LOAD_IM2COL{};
}
else {
return cute::SM100_TMA_2SM_LOAD_IM2COL_MULTICAST{};
}
}
else {
return cute::SM100_TMA_2SM_LOAD_IM2COL_MULTICAST{};
}
}
else if constexpr (size(atom_thr_id) == 1) {
if constexpr (!IsDynamicCluster) {
return detail::sm90_cluster_shape_to_im2col_tma_atom(cute::size<1>(cluster_shape_mnk));
}
else {
// In the case of dynamic cluster, multicast decision is not known at compile time.
// A multicast instruction is forced by passing a cute::Int<2>{} to this helper.
return detail::sm90_cluster_shape_to_im2col_tma_atom(cute::Int<2>{});
}
}
else {
static_assert(cutlass::detail::dependent_false<ClusterShapeMNK>,
"Unsupported Configuration for SM100 TMA");
}
}
template <class ClusterShapeMNK, class AtomThrId>
constexpr auto
sm100_cluster_shape_to_im2col_tma_atom_B(ClusterShapeMNK cluster_shape_mnk, AtomThrId atom_thr_id) {
static_assert(cute::rank(cluster_shape_mnk) == 3);
constexpr bool IsDynamicCluster = not cute::is_static_v<ClusterShapeMNK>;
if constexpr (cute::size(atom_thr_id) == 2) {
if constexpr (!IsDynamicCluster) {
static_assert(cute::size<0>(cluster_shape_mnk) % 2 == 0, "Cluster shape not divisible by MMA size");
if constexpr (cute::size<0>(cluster_shape_mnk) == 2) {
return cute::SM100_TMA_2SM_LOAD_IM2COL{};
}
else {
return cute::SM100_TMA_2SM_LOAD_IM2COL_MULTICAST{};
}
}
else {
return cute::SM100_TMA_2SM_LOAD_IM2COL_MULTICAST{};
}
} else if constexpr (size(atom_thr_id) == 1) {
if constexpr (!IsDynamicCluster) {
return detail::sm90_cluster_shape_to_im2col_tma_atom(cute::size<0>(cluster_shape_mnk));
}
else {
// In the case of dynamic cluster, multicast decision is not known at compile time.
// A multicast instruction is forced by passing a cute::Int<2>{} to this helper.
return detail::sm90_cluster_shape_to_im2col_tma_atom(cute::Int<2>{});
}
}
else {
static_assert(cutlass::detail::dependent_false<ClusterShapeMNK>,
"Unsupported Configuration for SM100 TMA");
}
}
template<
class ElementA,
class ElementB,
class ElementAccumulator,
class TileShape_MNK,
class ClusterShape_MNK,
UMMA::Major UmmaMajorA,
UMMA::Major UmmaMajorB,
class KernelScheduleType
>
constexpr auto
sm100_make_tiled_mma() {
// MMA_2SM requested
if constexpr (cute::is_same_v<KernelScheduleType, KernelImplicitTmaWarpSpecialized2SmSm100>) {
return cutlass::gemm::collective::detail::sm100_make_2sm_trivial_tiled_mma<
ElementA, ElementB, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB>();
}
// MMA_1SM requested
else if constexpr (cute::is_same_v<KernelScheduleType, KernelImplicitTmaWarpSpecialized1SmSm100>) {
return cutlass::gemm::collective::detail::sm100_make_1sm_trivial_tiled_mma<
ElementA, ElementB, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB>();
}
// Auto scheduling requested
else if constexpr (cute::is_same_v<KernelScheduleType, KernelScheduleAuto>) {
// Static cluster
if constexpr (cute::is_static_v<ClusterShape_MNK>) {
// For MMA_2SM we need a cluster shape that is multiple of 2x1
// and only M=128 and M=256 are supported, otherwise, fall back to MMA_1SM
if constexpr (cute::size<0>(ClusterShape_MNK{}) % 2 == 0 &&
cute::size<0>(TileShape_MNK{}) % 128 == 0) {
return cutlass::gemm::collective::detail::sm100_make_2sm_trivial_tiled_mma<
ElementA, ElementB, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB>();
}
else {
return cutlass::gemm::collective::detail::sm100_make_1sm_trivial_tiled_mma<
ElementA, ElementB, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB>();
}
// Dynamic cluster shape means we cannot assume we can use 2SM MMA
}
else {
return cutlass::gemm::collective::detail::sm100_make_1sm_trivial_tiled_mma<
ElementA, ElementB, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB>();
}
}
else {
static_assert(cutlass::detail::dependent_false<ElementA>,
"Unsupported policy for SM100 collective builder.");
}
}
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::conv::collective::detail
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -0,0 +1,225 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2025 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.
*
**************************************************************************************************/
//
//
#pragma once
#include "cutlass/conv/collective/builders/sm100_common.inl"
#include "cutlass/conv/collective/builders/sm90_gmma_builder.inl"
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass::conv::collective {
using namespace cute;
/////////////////////////////////////////////////////////////////////////////////////////////////
template <
conv::Operator ConvOp,
class ElementA,
class GmemLayoutA,
int AlignmentA,
class ElementB,
class GmemLayoutB,
int AlignmentB,
class ElementAccumulator,
class TileShape_MNKL, // (MmaAtomShapeM, MmaAtomShapeN, TileK, optional: TileL)
class ClusterShape_MNK, // Static cluster shape or dynamic (int, int, _1)
class StageCountType,
class KernelScheduleType
>
struct CollectiveBuilder<
arch::Sm100,
arch::OpClassTensorOp,
ConvOp,
ElementA,
GmemLayoutA,
AlignmentA,
ElementB,
GmemLayoutB,
AlignmentB,
ElementAccumulator,
TileShape_MNKL,
ClusterShape_MNK,
StageCountType,
KernelScheduleType,
cute::enable_if_t<
(cute::is_same_v<KernelScheduleType, KernelImplicitTmaWarpSpecialized1SmSm100> ||
cute::is_same_v<KernelScheduleType, KernelImplicitTmaWarpSpecialized2SmSm100> ||
cute::is_same_v<KernelScheduleType, KernelStridedDgradTmaWs1SmSm100> ||
cute::is_same_v<KernelScheduleType, KernelStridedDgradTmaWs2SmSm100> ||
cute::is_same_v<KernelScheduleType, KernelScheduleAuto>) &&
((sizeof(ElementA) * AlignmentA) % cutlass::gemm::collective::detail::tma_alignment_bytes == 0) &&
((sizeof(ElementB) * AlignmentB) % cutlass::gemm::collective::detail::tma_alignment_bytes == 0)>> {
private:
// For fprop, majorA = K, major B = K;
// For wgrad, majorA = MN, major B = MN;
// For dgrad, majorA = K, major B = MN;
static constexpr cute::UMMA::Major UmmaMajorA =
(ConvOp == conv::Operator::kWgrad) ? cute::UMMA::Major::MN : cute::UMMA::Major::K;
static constexpr cute::UMMA::Major UmmaMajorB =
(ConvOp == conv::Operator::kFprop) ? cute::UMMA::Major::K : cute::UMMA::Major::MN;
// For fp32 types, map to tf32 MMA value type
using ElementAMma = cute::conditional_t<cute::is_same_v<ElementA, float>, tfloat32_t, ElementA>;
using ElementBMma = cute::conditional_t<cute::is_same_v<ElementB, float>, tfloat32_t, ElementB>;
using TileShape_MNK = decltype(cute::take<0,3>(TileShape_MNKL{})); // (MmaAtomShapeM, MmaAtomShapeN, TileK)
static constexpr auto
get_tiled_mma_schedule() {
if constexpr (cute::is_same_v<KernelScheduleType, KernelStridedDgradTmaWs1SmSm100>) {
return KernelImplicitTmaWarpSpecialized1SmSm100{};
}
else if constexpr (cute::is_same_v<KernelScheduleType, KernelStridedDgradTmaWs2SmSm100>) {
return KernelImplicitTmaWarpSpecialized2SmSm100{};
}
else {
return KernelScheduleType{};
}
}
using TiledMmaSchedule = decltype(get_tiled_mma_schedule());
using TiledMma = decltype(detail::sm100_make_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
TileShape_MNK, ClusterShape_MNK,
UmmaMajorA, UmmaMajorB, TiledMmaSchedule>());
using AtomThrID = typename TiledMma::AtomThrID;
// ((MMA_TILE_M,MMA_TILE_K), MMA_M, MMA_K)
using MmaShapeA_MK = decltype(partition_shape_A(TiledMma{}, make_shape(cute::size<0>(TileShape_MNK{}),
cute::size<2>(TileShape_MNK{}))));
// ((MMA_TILE_N,MMA_TILE_K), MMA_N, MMA_K)
using MmaShapeB_NK = decltype(partition_shape_B(TiledMma{}, make_shape(cute::size<1>(TileShape_MNK{}),
cute::size<2>(TileShape_MNK{}))));
static constexpr auto
get_tma_atom_A() {
if constexpr (cute::is_same_v<KernelScheduleType,KernelStridedDgradTmaWs1SmSm100> ||
cute::is_same_v<KernelScheduleType,KernelStridedDgradTmaWs2SmSm100>) {
static_assert(ConvOp == conv::Operator::kDgrad, "Operator+Schedule mismatch");
return cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_A(ClusterShape_MNK{}, AtomThrID{});
}
else if constexpr (ConvOp == conv::Operator::kWgrad) {
return cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_A(ClusterShape_MNK{}, AtomThrID{});
}
else {
return cutlass::conv::collective::detail::sm100_cluster_shape_to_im2col_tma_atom_A(ClusterShape_MNK{}, AtomThrID{});
}
}
static constexpr auto
get_tma_atom_B() {
if constexpr (cute::is_same_v<KernelScheduleType,KernelStridedDgradTmaWs1SmSm100> ||
cute::is_same_v<KernelScheduleType,KernelStridedDgradTmaWs2SmSm100>) {
static_assert(ConvOp == conv::Operator::kDgrad, "Operator+Schedule mismatch");
return cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_B(ClusterShape_MNK{}, AtomThrID{});
}
else if constexpr (ConvOp == conv::Operator::kWgrad) {
return cutlass::conv::collective::detail::sm100_cluster_shape_to_im2col_tma_atom_B(ClusterShape_MNK{}, AtomThrID{});
}
else {
return cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_B(ClusterShape_MNK{}, AtomThrID{});
}
}
// For wgrad kernel, tensor A uses tma tiled mode and tensor B uses tma im2col mode.
using GmemTiledCopyA = decltype(get_tma_atom_A());
using GmemTiledCopyB = decltype(get_tma_atom_B());
using BlockTileA_M = decltype(cute::size<0,0>(MmaShapeA_MK{}) * cute::size<1>(MmaShapeA_MK{}));
using BlockTileA_K = decltype(cute::size<0,1>(MmaShapeA_MK{}) * cute::size<2>(MmaShapeA_MK{}));
using SmemLayoutAtomA = decltype(cutlass::gemm::collective::detail::sm100_smem_selector<
UmmaMajorA, ElementAMma, BlockTileA_M, BlockTileA_K>());
using BlockTileB_N = decltype(cute::size<0,0>(MmaShapeB_NK{}) * cute::size<1>(MmaShapeB_NK{}));
using BlockTileB_K = decltype(cute::size<0,1>(MmaShapeB_NK{}) * cute::size<2>(MmaShapeB_NK{}));
using SmemLayoutAtomB = decltype(cutlass::gemm::collective::detail::sm100_smem_selector<
UmmaMajorB, ElementBMma, BlockTileB_N, BlockTileB_K>());
// Calculate SMEM matrix A and B buffers' pipeline stages
static constexpr uint32_t AccumulatorPipelineStageCount = 2;
static constexpr uint32_t SchedulerPipelineStageCount = 2;
static constexpr uint32_t CLCResponseSize = 16;
// AccumulatorPipeline = PipelineUmmaAsync
static constexpr auto AccumulatorPipelineStorage = sizeof(typename cutlass::PipelineUmmaAsync<AccumulatorPipelineStageCount>::SharedStorage);
// CLCPipeline = PipelineCLCFetchAsync
static constexpr auto CLCPipelineStorage = sizeof(typename cutlass::PipelineCLCFetchAsync<SchedulerPipelineStageCount, ClusterShape_MNK>::SharedStorage);
// LoadOrderBarrier = OrderedSequenceBarrier<1,2>
static constexpr auto LoadOrderBarrierStorage = sizeof(typename cutlass::OrderedSequenceBarrier<1,2>::SharedStorage);
// CLC (scheduler) response
static constexpr auto CLCResponseStorage = SchedulerPipelineStageCount * CLCResponseSize;
// CLC Throttle pipeline storage
static constexpr auto CLCThrottlePipelineStorage = sizeof(typename cutlass::PipelineAsync<SchedulerPipelineStageCount>::SharedStorage);
// Tmem dealloc
static constexpr auto TmemDeallocStorage = sizeof(cutlass::arch::ClusterBarrier);
// Tmem ptr storage
static constexpr auto TmemBasePtrsStorage = SchedulerPipelineStageCount * sizeof(uint32_t);
// Smem usage that's not part of CollectiveEpilogue::SharedStorage & CollectiveMainloop::SharedStorage
static constexpr auto KernelSmemCarveout = static_cast<int>( AccumulatorPipelineStorage +
CLCPipelineStorage +
LoadOrderBarrierStorage +
TmemDeallocStorage +
CLCThrottlePipelineStorage +
CLCResponseStorage +
TmemBasePtrsStorage);
// Reduce SMEM capacity available for buffers considering barrier allocations.
static constexpr int Sm100ReducedSmemCapacityBytes = cutlass::gemm::collective::detail::sm100_smem_capacity_bytes - KernelSmemCarveout;
using SmemTileShape = cute::Shape<BlockTileA_M, BlockTileB_N, BlockTileA_K>;
static constexpr int PipelineStages = detail::compute_stage_count_or_override<
Sm100ReducedSmemCapacityBytes, ElementAMma, ElementBMma, SmemTileShape>(StageCountType{});
constexpr static int NumSpatialDimensions = detail::gmem_layout_tags_to_spatial_dims<GmemLayoutA, GmemLayoutB>();
using DispatchPolicy = cutlass::conv::MainloopSm100TmaUmmaWarpSpecializedImplicitGemm<
ConvOp, PipelineStages, NumSpatialDimensions, ClusterShape_MNK>;
public:
using CollectiveOp = cutlass::conv::collective::CollectiveConv<
DispatchPolicy,
TileShape_MNKL,
ElementA,
ElementB,
TiledMma,
detail::Sm100ImplicitGemmTileTraits<GmemTiledCopyA, SmemLayoutAtomA>,
detail::Sm100ImplicitGemmTileTraits<GmemTiledCopyB, SmemLayoutAtomB>
>;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::conv::collective
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -90,4 +90,5 @@ struct CollectiveBuilder {
/////////////////////////////////////////////////////////////////////////////////////////////////
#include "builders/sm90_gmma_builder.inl"
#include "builders/sm100_umma_builder.inl"
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -59,4 +59,5 @@ struct CollectiveConv {
/////////////////////////////////////////////////////////////////////////////////////////////////
#include "sm90_implicit_gemm_gmma_ss_warpspecialized.hpp"
#include "sm100_implicit_gemm_umma_warpspecialized.hpp"
/////////////////////////////////////////////////////////////////////////////////////////////////
+19 -2
View File
@@ -167,6 +167,20 @@ sm90_dispatch_policy_to_stride_B() {
}
}
template <class DispatchPolicy>
constexpr auto
sm100_dispatch_policy_to_stride_A() {
return sm90_dispatch_policy_to_stride_A<DispatchPolicy>();
}
template <class DispatchPolicy>
constexpr auto
sm100_dispatch_policy_to_stride_B() {
return sm90_dispatch_policy_to_stride_B<DispatchPolicy>();
}
/////////////////////////////////////////////////////////////////////////////////////////////////
// Compute the lower/near corner, returning it as a cute::array in [W,H,D] order
@@ -247,8 +261,11 @@ compute_lower_srt(ConvProblemShape<ConvOp, NumSpatialDimensions> const& problem_
}
template <class CopyOp> struct is_im2col_load { static constexpr bool value = false; };
template <> struct is_im2col_load<SM90_TMA_LOAD_IM2COL > { static constexpr bool value = true; };
template <> struct is_im2col_load<SM90_TMA_LOAD_IM2COL_MULTICAST> { static constexpr bool value = true; };
template <> struct is_im2col_load<cute::SM90_TMA_LOAD_IM2COL > { static constexpr bool value = true; };
template <> struct is_im2col_load<cute::SM90_TMA_LOAD_IM2COL_MULTICAST> { static constexpr bool value = true; };
template <> struct is_im2col_load<cute::SM100_TMA_2SM_LOAD_IM2COL > { static constexpr bool value = true; };
template <> struct is_im2col_load<cute::SM100_TMA_2SM_LOAD_IM2COL_MULTICAST> { static constexpr bool value = true; };
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::conv::collective::detail
@@ -0,0 +1,899 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2025 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.
*
**************************************************************************************************/
//
//
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/gemm/dispatch_policy.hpp"
#include "cutlass/pipeline/pipeline.hpp"
#include "cutlass/gemm/gemm.h"
#include "cutlass/detail/cluster.hpp"
#include "cutlass/conv/detail.hpp"
#include "cute/algorithm/functional.hpp"
#include "cute/arch/cluster_sm90.hpp"
#include "cute/atom/mma_atom.hpp"
#include "cute/algorithm/gemm.hpp"
#include "cute/tensor_predicate.hpp"
#include "cute/numeric/arithmetic_tuple.hpp"
#include "cutlass/trace.h"
#if (! defined(__CUDA_ARCH__)) && (CUTLASS_DEBUG_TRACE_LEVEL > 0)
# include <sstream>
#endif
/////////////////////////////////////////////////////////////////////////////////////////////////
namespace cutlass::conv::collective {
using namespace cute;
/////////////////////////////////////////////////////////////////////////////////////////////////
// WarpSpecialized Mainloop
// Both DMA Load and MMA methods of this class must be run by a single thread that's picked by elect_one
template <
conv::Operator ConvOp,
int Stages,
int NumSpatialDims,
class ClusterShape, // Static cluster shape or dynamic (int, int, _1)
class TileShapeMNKL_, // (MmaAtomShapeM, MmaAtomShapeN, TileK, optional: TileL)
class ElementA_,
class ElementB_,
class TiledMma_,
class TileTraitsA_,
class TileTraitsB_>
struct CollectiveConv<
MainloopSm100TmaUmmaWarpSpecializedImplicitGemm<
ConvOp, Stages, NumSpatialDims, ClusterShape>,
TileShapeMNKL_,
ElementA_,
ElementB_,
TiledMma_,
TileTraitsA_,
TileTraitsB_>
{
//
// Type Aliases
//
using DispatchPolicy = MainloopSm100TmaUmmaWarpSpecializedImplicitGemm<
ConvOp, Stages, NumSpatialDims, ClusterShape>;
using TileShape = decltype(cute::take<0,3>(TileShapeMNKL_{})); // (MmaAtomShapeM, MmaAtomShapeN, TileK)
using ElementA = ElementA_;
using ElementB = ElementB_;
using TiledMma = TiledMma_;
using ElementAccumulator = typename TiledMma::ValTypeC;
using GmemTiledCopyA = typename TileTraitsA_::GmemTiledCopy;
using GmemTiledCopyB = typename TileTraitsB_::GmemTiledCopy;
using SmemLayoutAtomA = typename TileTraitsA_::SmemLayoutAtom;
using SmemLayoutAtomB = typename TileTraitsB_::SmemLayoutAtom;
using ArchTag = typename DispatchPolicy::ArchTag;
static constexpr int NumSpatialDimensions = DispatchPolicy::NumSpatialDimensions;
static constexpr int NumTensorDimensions = NumSpatialDimensions + 2;
// deducde the kernel facing stride tuple types based on the dispatch policy (spatial dim, algo, etc.)
using StrideA = decltype(detail::sm100_dispatch_policy_to_stride_A<DispatchPolicy>());
using StrideB = decltype(detail::sm100_dispatch_policy_to_stride_B<DispatchPolicy>());
static constexpr bool IsDynamicCluster = not cute::is_static_v<ClusterShape>;
static constexpr bool ConvertF32toTF32A = cute::is_same_v<float, ElementA>;
static constexpr bool ConvertF32toTF32B = cute::is_same_v<float, ElementB>;
using TmaInternalElementA = cute::conditional_t<ConvertF32toTF32A, tfloat32_t, cute::uint_bit_t<cute::sizeof_bits_v<ElementA>>>;
using TmaInternalElementB = cute::conditional_t<ConvertF32toTF32B, tfloat32_t, cute::uint_bit_t<cute::sizeof_bits_v<ElementB>>>;
using ElementAMma = cute::conditional_t<cute::is_same_v<ElementA, float>, tfloat32_t, ElementA>;
using ElementBMma = cute::conditional_t<cute::is_same_v<ElementB, float>, tfloat32_t, ElementB>;
// Determine MMA type: MMA_1SM vs MMA_2SM
using AtomThrShapeMNK = Shape<decltype(shape<0>(typename TiledMma_::ThrLayoutVMNK{})), _1, _1>;
using MainloopPipeline = cutlass::PipelineTmaUmmaAsync<
DispatchPolicy::Stages,
ClusterShape,
AtomThrShapeMNK>;
using MainloopPipelineState = typename MainloopPipeline::PipelineState;
using ProblemShape = ConvProblemShape<ConvOp, NumSpatialDimensions>;
CUTE_STATIC_ASSERT_V(evenly_divides(shape<0>(TileShape{}), tile_size<0>(TiledMma{})), "TileShape_M should be evenly divided by TiledMma_M");
CUTE_STATIC_ASSERT_V(evenly_divides(shape<1>(TileShape{}), tile_size<1>(TiledMma{})) || (ConvOp == conv::Operator::kWgrad), "TileShape_N should be evenly divided by TiledMma_N");
using CtaShape_MNK = decltype(shape_div(TileShape{}, AtomThrShapeMNK{}));
// Define A and B block shapes for reduced size TMA_LOADs
using MmaShapeA_MK = decltype(partition_shape_A(TiledMma{}, make_shape(size<0>(TileShape{}), size<2>(TileShape{}))));
using MmaShapeB_NK = decltype(partition_shape_B(TiledMma{}, make_shape(size<1>(TileShape{}), size<2>(TileShape{}))));
static_assert(rank(SmemLayoutAtomA{}) == 2, "SmemLayoutAtom must be rank 2 (M/N, K)");
static_assert(((size<0,0>(MmaShapeA_MK{}) * size<1>(MmaShapeA_MK{})) % size<0>(SmemLayoutAtomA{})) == 0,
"SmemLayoutAtom must evenly divide tile shape.");
static_assert(((size<0,1>(MmaShapeA_MK{}) * size<2>(MmaShapeA_MK{})) % size<1>(SmemLayoutAtomA{})) == 0,
"SmemLayoutAtom must evenly divide tile shape.");
static_assert(rank(SmemLayoutAtomB{}) == 2, "SmemLayoutAtom must be rank 2 (M/N, K)");
static_assert(((size<0,0>(MmaShapeB_NK{}) * size<1>(MmaShapeB_NK{})) % size<0>(SmemLayoutAtomB{})) == 0,
"SmemLayoutAtom must evenly divide tile shape.");
static_assert(((size<0,1>(MmaShapeB_NK{}) * size<2>(MmaShapeB_NK{})) % size<1>(SmemLayoutAtomB{})) == 0,
"SmemLayoutAtom must evenly divide tile shape.");
// Tile along K mode first before tiling over MN. PIPE mode last as usual.
// This maximizes TMA boxes due to better smem-K vectorization, reducing total issued TMAs.
using SmemLayoutA = decltype(UMMA::tile_to_mma_shape(
SmemLayoutAtomA{},
append(MmaShapeA_MK{}, Int<DispatchPolicy::Stages>{}),
Step<_2,_1,_3>{}));
using SmemLayoutB = decltype(UMMA::tile_to_mma_shape(
SmemLayoutAtomB{},
append(MmaShapeB_NK{}, Int<DispatchPolicy::Stages>{}),
Step<_2,_1,_3>{}));
static_assert(DispatchPolicy::Stages >= 2, "Specialization requires Stages set to value 1 or more.");
static_assert(cute::is_base_of<cute::UMMA::DescriptorIterator, typename TiledMma::FrgTypeA>::value &&
cute::is_base_of<cute::UMMA::DescriptorIterator, typename TiledMma::FrgTypeB>::value,
"MMA atom must source both A and B operand from smem_desc for this mainloop.");
static constexpr bool is_im2col_A = detail::is_im2col_load<GmemTiledCopyA>::value;
static constexpr bool is_im2col_B = detail::is_im2col_load<GmemTiledCopyB>::value;
static constexpr bool is_strided_dgrad = ConvOp == conv::Operator::kDgrad && not is_im2col_A && not is_im2col_B;
static constexpr int TileShapeMNKLRank = rank(TileShapeMNKL_{});
// If rank > 3, TileL exists and it is GroupsPerTile. The kernel is grouped conv now.
static constexpr bool is_grouped_wgrad = ConvOp == conv::Operator::kWgrad && TileShapeMNKLRank > 3;
struct SharedStorage {
struct TensorStorage : cute::aligned_struct<128, _0> {
cute::array_aligned<typename TiledMma::ValTypeA, cute::cosize_v<SmemLayoutA>> smem_A;
cute::array_aligned<typename TiledMma::ValTypeB, cute::cosize_v<SmemLayoutB>> smem_B;
} tensors;
using PipelineStorage = typename MainloopPipeline::SharedStorage;
PipelineStorage pipeline;
};
using TensorStorage = typename SharedStorage::TensorStorage;
using PipelineStorage = typename SharedStorage::PipelineStorage;
// Only one thread issues the TMA and updates the barriers in a 2SM MMA, adjust bytes accordingly
static constexpr uint32_t TmaTransactionBytes =
size(AtomThrShapeMNK{}) * (size<0>(SmemLayoutA{}) * size<1>(SmemLayoutA{}) * size<2>(SmemLayoutA{}) * static_cast<uint32_t>(sizeof(ElementA))) +
size(AtomThrShapeMNK{}) * (size<0>(SmemLayoutB{}) * size<1>(SmemLayoutB{}) * size<2>(SmemLayoutB{}) * static_cast<uint32_t>(sizeof(ElementB)));
// Host side kernel arguments
struct Arguments {
ElementA const* ptr_A{nullptr};
ElementB const* ptr_B{nullptr};
};
private:
// Note that for fprop and non-strided dgrad kernel, the tma load mode is im2col for tensor A and tiled for
// tensor B while for wgrad kernel, the tma load mode is tiled for tensor A and im2col for tensor
// B since operand A, B is swapped.
// For strided dgrad A and B are both tma tiled and not im2col
template <class TensorA, class ClusterShapeVMNK>
static constexpr auto
get_tma_load_a_instance(
TensorA const& tensor_a,
ProblemShape const& problem_shape,
ClusterShapeVMNK const& cluster_shape_vmnk) {
if constexpr (is_im2col_A) {
// compute the upper and lower corners based on the conv padding
auto lower_corner_whd = detail::compute_lower_corner_whd(problem_shape);
auto upper_corner_whd = detail::compute_upper_corner_whd(problem_shape);
auto lower_srt = detail::compute_lower_srt(problem_shape);
// gbasis strides for dgrad kernel need to be negated
cute::array<int32_t, NumSpatialDimensions> stride_srt{};
for (int i = 0; i < NumSpatialDimensions; ++i) {
stride_srt[i] = ConvOp == conv::Operator::kDgrad ?
-problem_shape.dilation[NumSpatialDimensions-1-i] :
problem_shape.dilation[NumSpatialDimensions-1-i];
}
return make_im2col_tma_atom_A_sm100(
GmemTiledCopyA{},
tensor_a,
SmemLayoutA{}(_,_,_,cute::Int<0>{}),
TileShape{},
TiledMma{},
cluster_shape_vmnk,
shape(lower_corner_whd),
shape(upper_corner_whd),
cute::reverse(shape(problem_shape.lower_padding)),
cute::reverse(shape(problem_shape.upper_padding)),
cute::reverse(shape(problem_shape.traversal_stride)),
shape(lower_srt),
shape(stride_srt));
}
// TMA tiled mode for tensor A in wgrad and strided dgrad
else {
return make_tma_atom_A_sm100<TmaInternalElementA>(
GmemTiledCopyA{},
tensor_a,
SmemLayoutA{}(_,_,_,cute::Int<0>{}),
TileShape{},
TiledMma{},
cluster_shape_vmnk);
}
}
template <class TensorB, class ClusterShapeVMNK>
static constexpr auto
get_tma_load_b_instance(
TensorB const& tensor_b,
ProblemShape const& problem_shape,
ClusterShapeVMNK const& cluster_shape_vmnk) {
if constexpr (is_im2col_B) {
// compute the upper and lower corners based on the conv padding
auto lower_corner_whd = detail::compute_lower_corner_whd(problem_shape);
auto upper_corner_whd = detail::compute_upper_corner_whd(problem_shape);
auto lower_srt = detail::compute_lower_srt(problem_shape);
return make_im2col_tma_atom_B_sm100(
GmemTiledCopyB{},
tensor_b,
SmemLayoutB{}(_,_,_,cute::Int<0>{}),
TileShape{},
TiledMma{},
cluster_shape_vmnk,
shape(lower_corner_whd),
shape(upper_corner_whd),
cute::reverse(shape(problem_shape.lower_padding)),
cute::reverse(shape(problem_shape.upper_padding)),
cute::reverse(shape(problem_shape.traversal_stride)),
shape(lower_srt),
cute::reverse(shape(problem_shape.dilation)));
}
else {
return make_tma_atom_B_sm100<TmaInternalElementB>(
GmemTiledCopyB{},
tensor_b,
SmemLayoutB{}(_,_,_,cute::Int<0>{}),
TileShape{},
TiledMma{},
cluster_shape_vmnk);
}
}
public:
// Performs im2col transformations on the input of type ConvProblemShape
static constexpr auto
get_problem_shape_MNKL(ProblemShape const& problem_shape) {
if constexpr (is_im2col_A || is_im2col_B) {
// transformation + im2col linearization
return cutlass::conv::detail::get_linearized_problem_shape_MNKL(problem_shape);
}
else {
// transformation
return cutlass::conv::detail::get_transformed_problem_shape_MNKL(problem_shape);
}
}
// Device-side kernel params
//
// Arguments has the untransformed problem shape from the user.
// Params will have the transformed problem shape.
struct Params {
using _Submode = decltype(take<0,NumTensorDimensions-1>(typename ProblemShape::TensorExtent{}));
using ClusterLayout_VMNK = decltype(tiled_divide(make_layout(conditional_return<IsDynamicCluster>(make_shape(uint32_t(0), uint32_t(0), Int<1>{}), ClusterShape{})),
make_tile(typename TiledMma::AtomThrID{})));
// Assumption: StrideA is congruent with Problem_MK
// Select TMA load type according to convolution operator.
using TensorShapeA = cute::conditional_t<ConvOp == conv::Operator::kWgrad,
decltype(repeat_like(StrideA{}, int32_t(0))),
decltype(make_shape(_Submode{}, int32_t(0)))>;
using TensorShapeB = cute::conditional_t<ConvOp == conv::Operator::kWgrad,
decltype(make_shape(int32_t(0), _Submode{})),
decltype(repeat_like(StrideB{}, int32_t(0)))>;
using TMA_A = decltype(get_tma_load_a_instance(
make_tensor(
make_gmem_ptr(recast_ptr<TmaInternalElementA>(nullptr)),
make_layout(TensorShapeA{}, StrideA{})),
ConvProblemShape<ConvOp, NumSpatialDimensions>{},
ClusterLayout_VMNK{}));
using TMA_B = decltype(get_tma_load_b_instance(
make_tensor(
make_gmem_ptr(recast_ptr<TmaInternalElementB>(nullptr)),
make_layout(TensorShapeB{}, StrideB{})),
ConvProblemShape<ConvOp, NumSpatialDimensions>{},
ClusterLayout_VMNK{}));
// Members
TMA_A tma_load_a;
TMA_B tma_load_b;
TMA_A tma_load_a_fallback;
TMA_B tma_load_b_fallback;
dim3 cluster_shape_fallback;
};
//
// Constructor
//
CUTLASS_DEVICE
CollectiveConv(Params const& params) {
if constexpr (IsDynamicCluster) {
dim3 cs = cute::cluster_shape();
const bool is_fallback_cluster = (cs.x == params.cluster_shape_fallback.x && cs.y == params.cluster_shape_fallback.y);
observed_tma_load_a_ = is_fallback_cluster ? &params.tma_load_a_fallback : &params.tma_load_a;
observed_tma_load_b_ = is_fallback_cluster ? &params.tma_load_b_fallback : &params.tma_load_b;
}
else {
observed_tma_load_a_ = &params.tma_load_a;
observed_tma_load_b_ = &params.tma_load_b;
}
}
//
// Methods
//
static constexpr Params
to_underlying_arguments(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cutlass::KernelHardwareInfo const& hw_info = cutlass::KernelHardwareInfo{}) {
(void) workspace;
// from the flat problem shape arrays of ConvProblemShape<N>, create a rank-3 MNK problem shape tuple
// tma desc creation depends on the original untransformed domain.
// A extents.
auto shape_A_orig = problem_shape.get_shape_A();
// B extents.
auto shape_B_orig = problem_shape.get_shape_B();
// Fill inferred cute strides from flat stride arrays
auto dA = make_cute_packed_stride(StrideA{}, problem_shape.stride_A, ConvOp);
auto dB = make_cute_packed_stride(StrideB{}, problem_shape.stride_B, ConvOp);
auto ptr_A = recast_ptr<TmaInternalElementA>(args.ptr_A);
auto ptr_B = recast_ptr<TmaInternalElementB>(args.ptr_B);
Tensor tensor_a = make_tensor(make_gmem_ptr(ptr_A), make_layout(shape_A_orig, dA));
Tensor tensor_b = make_tensor(make_gmem_ptr(ptr_B), make_layout(shape_B_orig, dB));
auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape);
// Cluster layout for TMA construction
auto cluster_layout_vmnk = tiled_divide(make_layout(cluster_shape), make_tile(typename TiledMma::AtomThrID{}));
auto cluster_shape_fallback = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape_fallback);
// Cluster layout for TMA construction
auto cluster_layout_vmnk_fallback = tiled_divide(make_layout(cluster_shape_fallback), make_tile(typename TiledMma::AtomThrID{}));
auto tma_load_a = get_tma_load_a_instance(tensor_a, problem_shape, cluster_layout_vmnk);
auto tma_load_b = get_tma_load_b_instance(tensor_b, problem_shape, cluster_layout_vmnk);
auto tma_load_a_fallback = get_tma_load_a_instance(tensor_a, problem_shape, cluster_layout_vmnk_fallback);
auto tma_load_b_fallback = get_tma_load_b_instance(tensor_b, problem_shape, cluster_layout_vmnk_fallback);
static_assert(size(typename decltype(tma_load_a)::ThrID{}) == size(AtomThrShapeMNK{}));
static_assert(size(typename decltype(tma_load_b)::ThrID{}) == size(AtomThrShapeMNK{}));
return {
tma_load_a,
tma_load_b,
tma_load_a_fallback,
tma_load_b_fallback,
hw_info.cluster_shape_fallback
};
}
template<class ProblemShape>
static bool
can_implement(
ProblemShape const& problem_shape,
Arguments const& args) {
// Activation and Filter channel mode extents much match
bool implementable = true;
// channel mode is major
{
const bool check = problem_shape.stride_A[NumTensorDimensions-1] == 1;
#if (! defined(__CUDA_ARCH__)) && (CUTLASS_DEBUG_TRACE_LEVEL > 0)
if (not check) {
const auto offending_stride =
problem_shape.stride_A[NumTensorDimensions-1];
std::ostringstream os;
os << "CollectiveConv::can_implement: "
"problem_shape.stride_A[NumTensorDimensions-1 = "
<< (NumTensorDimensions-1) << "] = "
<< offending_stride << " != 1";
CUTLASS_TRACE_HOST( os.str() );
}
#endif
implementable &= check;
}
{
const bool check = problem_shape.stride_B[NumTensorDimensions-1] == 1;
#if (! defined(__CUDA_ARCH__)) && (CUTLASS_DEBUG_TRACE_LEVEL > 0)
if (not check) {
const auto offending_stride =
problem_shape.stride_B[NumTensorDimensions-1];
std::ostringstream os;
os << "CollectiveConv::can_implement: "
"problem_shape.stride_B[NumTensorDimensions-1 = "
<< (NumTensorDimensions-1) << "] = "
<< offending_stride << " != 1\n";
CUTLASS_TRACE_HOST( os.str() );
}
#endif
implementable &= check;
}
{
const auto & traversal_stride = problem_shape.traversal_stride;
for (auto stride: traversal_stride) {
implementable &= (stride >= 1 && stride <= 8);
}
}
if constexpr (ConvOp == conv::Operator::kDgrad && not is_strided_dgrad) {
const auto & traversal_stride = problem_shape.traversal_stride;
for (auto stride: traversal_stride) {
implementable &= (stride == 1);
}
}
constexpr int tma_alignment_bits = 128;
// A extents.
auto shape_A_orig = problem_shape.get_shape_A();
// B extents.
auto shape_B_orig = problem_shape.get_shape_B();
constexpr int min_tma_aligned_elements_A = tma_alignment_bits / cutlass::sizeof_bits<ElementA>::value;
{
const bool check = cutlass::detail::check_alignment<min_tma_aligned_elements_A>(shape_A_orig, StrideA{});
if (not check) {
CUTLASS_TRACE_HOST("A shape and/or strides have alignment issue.");
}
implementable &= check;
}
constexpr int min_tma_aligned_elements_B = tma_alignment_bits / cutlass::sizeof_bits<ElementB>::value;
{
const bool check = cutlass::detail::check_alignment<min_tma_aligned_elements_B>(shape_B_orig, StrideB{});
if (not check) {
CUTLASS_TRACE_HOST("B shape and/or strides have alignment issue.");
}
implementable &= check;
}
if (not implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Problem Size doesn't meet the minimum alignment requirements for TMA.\n");
return false;
}
if (is_im2col_A || is_im2col_B) {
// Check valid corner values for TMA_LOAD_IM2COL, signed int ranging from [-corner_limit, corner_limit - 1]
constexpr int32_t corner_limit = 1 << (16 / NumSpatialDimensions - 1);
auto lower_corner_whd = detail::compute_lower_corner_whd(problem_shape);
for (int i = 0; i < problem_shape.RankS; ++i) {
implementable = implementable && lower_corner_whd[i] >= -corner_limit && lower_corner_whd[i] <= (corner_limit - 1);
}
auto upper_corner_whd = detail::compute_upper_corner_whd(problem_shape);
for (int i = 0; i < problem_shape.RankS; ++i) {
implementable = implementable && upper_corner_whd[i] >= -corner_limit && upper_corner_whd[i] <= (corner_limit - 1);
}
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Padding values don't meet requirements for TMA LOAD IM2COL.\n");
return false;
}
}
if (is_im2col_A || is_im2col_B) {
// Check valid filter offsets for TMA_LOAD_IM2COL, unsigned int ranging from [0, offset_limit - 1]
constexpr int32_t offset_limit = 1 << (16 / NumSpatialDimensions);
auto flt_data = (ConvOp == conv::Operator::kWgrad) ? problem_shape.shape_C : problem_shape.shape_B;
for (int i = 0; i < problem_shape.RankS; ++i) {
// flt_data array contains [K, T, R, S, C], so pure filter [T, R, S] starts from the second position in the array
implementable = implementable && (flt_data[i+1] * problem_shape.dilation[i] >= 0)
&& (flt_data[i+1] * problem_shape.dilation[i] <= (offset_limit - 1));
}
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: tensor coordinate offset values don't meet requirements for TMA LOAD IM2COL.\n");
return false;
}
}
// Wgrad kernels don't support non-packed output strides, non-packed tensor A stride (linearized)
if constexpr (ConvOp == conv::Operator::kWgrad) {
const auto & input_shape = problem_shape.shape_A;
const auto & input_stride = problem_shape.stride_A;
implementable &= input_stride[ProblemShape::RankT - 1] == 1;
int input_shape_size = 1;
for (int i = ProblemShape::RankT - 2; i >= 0; --i) {
input_shape_size *= input_shape[i + 1];
implementable &= input_stride[i] == input_shape_size;
}
const auto & output_shape = problem_shape.shape_C;
const auto & output_stride = problem_shape.stride_C;
implementable &= output_stride[ProblemShape::RankT - 1] == 1;
int output_shape_size = 1;
for (int i = ProblemShape::RankT - 2; i >= 0; --i) {
output_shape_size *= output_shape[i + 1];
implementable &= output_stride[i] == output_shape_size;
}
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Wgrad kernels don't support non-packed output strides.\n");
return false;
}
}
// Conv kernels only support cross correlation mode currently.
{
implementable &= problem_shape.mode == cutlass::conv::Mode::kCrossCorrelation;
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Conv kernels only support cross correlation mode currently.\n");
return false;
}
}
// When groups > 1, it should be a Grouped Conv.
if (problem_shape.groups > 1) {
implementable &= TileShapeMNKLRank > 3;
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Only Grouped Conv can support groups > 1.\n");
return false;
}
}
// Only support Grouped Wgrad currently.
if constexpr (TileShapeMNKLRank > 3) {
implementable &= ConvOp == conv::Operator::kWgrad;
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Grouped Conv Only support Grouped Wgrad currently.\n");
return false;
}
}
// Grouped Wgrad channel check.
if constexpr (is_grouped_wgrad) {
int input_K = size<0>(problem_shape.get_shape_A());
int input_C = size<0>(problem_shape.get_shape_B());
implementable &= input_K == input_C;
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Grouped Conv's input K and input C do not match.\n");
return false;
}
int output_K = size<0>(problem_shape.get_shape_C());
int output_C = size<1,0>(problem_shape.get_shape_C());
implementable &= input_K == output_K;
implementable &= input_C == output_C * problem_shape.groups;
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Grouped Wgrad's input and output K,C and groups do not match\n");
return false;
}
constexpr int Tile_N = size<1>(TileShape{});
constexpr int GroupsPerTile = size<3>(TileShapeMNKL_{});
implementable &= Tile_N / GroupsPerTile == input_C / problem_shape.groups;
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Grouped Wgrad's Tile_N, GroupsPerTile and input_C, groups do not match.\n");
return false;
}
}
return true;
}
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
CUTLASS_DEVICE static void
prefetch_tma_descriptors(Params const& mainloop_params) {
if constexpr (IsDynamicCluster) {
dim3 cs = cute::cluster_shape();
const bool is_fallback_cluster = (cs.x == mainloop_params.cluster_shape_fallback.x && cs.y == mainloop_params.cluster_shape_fallback.y);
if (is_fallback_cluster) {
cute::prefetch_tma_descriptor(mainloop_params.tma_load_a_fallback.get_tma_descriptor());
cute::prefetch_tma_descriptor(mainloop_params.tma_load_b_fallback.get_tma_descriptor());
}
else {
cute::prefetch_tma_descriptor(mainloop_params.tma_load_a.get_tma_descriptor());
cute::prefetch_tma_descriptor(mainloop_params.tma_load_b.get_tma_descriptor());
}
}
else {
cute::prefetch_tma_descriptor(mainloop_params.tma_load_a.get_tma_descriptor());
cute::prefetch_tma_descriptor(mainloop_params.tma_load_b.get_tma_descriptor());
}
}
/// Construct A Single Stage's Accumulator Shape
CUTLASS_DEVICE auto
partition_accumulator_shape() {
auto acc_shape = partition_shape_C(TiledMma{}, take<0,2>(TileShape{})); // ((MMA_TILE_M,MMA_TILE_N),MMA_M,MMA_N)
return acc_shape;
}
/// Perform a collective-scoped matrix multiply-accumulate
/// Producer Perspective
template <
class GTensorA, class GTensorB,
class GTensorPartitionedA, class GTensorPartitionedB,
class STensorA, class STensorB,
class TileCoordMNKL,
class KTileIterator
>
CUTLASS_DEVICE auto
load(
Params const& params,
MainloopPipeline pipeline,
MainloopPipelineState mainloop_pipe_producer_state,
cute::tuple<GTensorA, GTensorB,
GTensorPartitionedA, GTensorPartitionedB,
STensorA, STensorB,
uint16_t, uint16_t> const& load_inputs,
TileCoordMNKL const& cta_coord_mnkl,
KTileIterator k_tile_iter, int k_tile_count) {
auto [unused_gA, unused_gB,
tAgA_mk, tBgB_nk, tAsA, tBsB,
mcast_mask_a, mcast_mask_b] = load_inputs;
// slice out the work coord from partitioned tensors
Tensor tAgA = tAgA_mk(_, get<0>(cta_coord_mnkl) / size(typename TiledMma::AtomThrID{}), _);
auto tensor_b_coord = get<1>(cta_coord_mnkl);
if constexpr (is_grouped_wgrad) {
// in grouped wgrad, tensor A = NZPQK, tensor B = NDHWC, tensor C = KTRSc, where C = G*c, c = channel_per_group = 8,16,32.
// CTA Tiling follows output tensor KTRSc. So cta_size_m = K/CTA_TILE_M. cta_size_n = T*R*S*ceil(c/CTA_TILE_N) = T*R*S*1 = T*R*S.
// tensor_a_coord = K_idx = cta_coord_m.
// tensor_b_coord = TRS_idx * C/CTA_TILE_N + C_idx = cta_coord_n * get<1,0>(shape(tBgB_nk) + cta_coord_m,
// because K == C and CTA_TILE_M == CTA_TILE_N => C_idx = K_idx = cta_coord_m.
tensor_b_coord = get<0>(cta_coord_mnkl) + get<1>(cta_coord_mnkl) * get<1,0>(shape(tBgB_nk));
}
Tensor tBgB = tBgB_nk(_, tensor_b_coord, _);
auto barrier_token = pipeline.producer_try_acquire(mainloop_pipe_producer_state);
// Issue the Mainloop loads
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
// LOCK mainloop_pipe_producer_state for _writing_
pipeline.producer_acquire(mainloop_pipe_producer_state, barrier_token);
using BarrierType = typename MainloopPipeline::ProducerBarrierType;
BarrierType* tma_barrier = pipeline.producer_get_barrier(mainloop_pipe_producer_state);
int write_stage = mainloop_pipe_producer_state.index();
++mainloop_pipe_producer_state;
barrier_token = pipeline.producer_try_acquire(mainloop_pipe_producer_state);
if constexpr (is_strided_dgrad) {
// construct gemm-k tile coord for gB
auto [conv_k, flt_coord, out_coord] = *k_tile_iter;
auto gemm_k_tile = prepend(flt_coord, conv_k); // (k,s,r,t)
// gA doesn't have a gemm-k (k,s,r,t) iterator mode because it's not an im2col tensor
auto offset_kqpzn = append(prepend(out_coord, _0{}),_0{}); // (k,q,p,z,n)
auto tAgA_offset = make_tensor(tAgA.data() + offset_kqpzn, tAgA.layout()); // (TMA, k)
if (cute::elect_one_sync()) {
copy(observed_tma_load_a_->with(*tma_barrier, mcast_mask_a), tAgA_offset(_,conv_k), tAsA(_,write_stage));
copy(observed_tma_load_b_->with(*tma_barrier, mcast_mask_b), tBgB(_,gemm_k_tile) , tBsB(_,write_stage));
}
}
else {
if (cute::elect_one_sync()) {
copy(observed_tma_load_a_->with(*tma_barrier, mcast_mask_a), tAgA(_,*k_tile_iter), tAsA(_,write_stage));
copy(observed_tma_load_b_->with(*tma_barrier, mcast_mask_b), tBgB(_,*k_tile_iter), tBsB(_,write_stage));
}
}
--k_tile_count;
++k_tile_iter;
}
return cute::make_tuple(mainloop_pipe_producer_state, k_tile_iter);
}
/// Set up the data needed by this collective for load.
/// Return tuple element contain
/// gA_mk - The tiled tma tensor for input A
/// gB_nk - The tiled tma tensor for input B
/// tAsA - partitioned smem tensor for A
/// tBsB - partitioned smem tensor for B
/// mcast_mask_a - tma multicast mask for A
/// mcast_mask_b - tma multicast mask for B
template <class ProblemShape_MNKL>
CUTLASS_DEVICE auto
load_init(
ProblemShape_MNKL const& problem_shape_MNKL,
Params const& params,
TensorStorage& shared_tensors) const {
using X = Underscore;
// Separate out problem shape for convenience
auto [M,N,K,L] = problem_shape_MNKL;
// Represent the full tensors -- get these from TMA
auto K_A = conditional_return<is_strided_dgrad>(get<0>(K), K);
Tensor mA_mk = observed_tma_load_a_->get_tma_tensor(make_shape(M, K_A));
Tensor mB_nk = observed_tma_load_b_->get_tma_tensor(make_shape(N, K));
// Tile the tensors and defer the slice
Tensor gA_mk = local_tile(mA_mk, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M, BLK_K, m, k)
Tensor gB_nk = local_tile(mB_nk, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N, BLK_K, n, k)
// Partition for this CTA
ThrMMA cta_mma = TiledMma{}.get_slice(blockIdx.x % size(typename TiledMma::AtomThrID{}));
Tensor tCgA_mk = cta_mma.partition_A(gA_mk); // (MMA, MMA_M, MMA_K, m, k)
Tensor tCgB_nk = cta_mma.partition_B(gB_nk); // (MMA, MMA_N, MMA_K, n, k)
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (MMA,MMA_M,MMA_K,PIPE)
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (MMA,MMA_N,MMA_K,PIPE)
auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, cute::cluster_shape());
Layout cta_layout_mnk = make_layout(cluster_shape);
Layout cta_layout_vmnk = tiled_divide(cta_layout_mnk, make_tile(typename TiledMma::AtomThrID{}));
int block_rank_in_cluster = cute::block_rank_in_cluster();
auto cta_coord_vmnk = cta_layout_vmnk.get_flat_coord(block_rank_in_cluster);
// Project the cta_layout for tma_a along the n-modes
auto [tAgA_mk, tAsA] = tma_partition(*observed_tma_load_a_,
get<2>(cta_coord_vmnk), make_layout(size<2>(cta_layout_vmnk)),
group_modes<0,3>(sA), group_modes<0,3>(tCgA_mk));
// Project the cta_layout for tma_b along the m-modes
auto [tBgB_nk, tBsB] = tma_partition(*observed_tma_load_b_,
get<1>(cta_coord_vmnk), make_layout(size<1>(cta_layout_vmnk)),
group_modes<0,3>(sB), group_modes<0,3>(tCgB_nk));
// TMA Multicast Masks
uint16_t mcast_mask_a = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk);
uint16_t mcast_mask_b = create_tma_multicast_mask<1>(cta_layout_vmnk, cta_coord_vmnk);
return cute::make_tuple(
gA_mk, gB_nk, // for scheduler
tAgA_mk, tBgB_nk, tAsA, tBsB, // for input tensor values
mcast_mask_a, mcast_mask_b); // multicast masks
}
/// Perform a Producer Epilogue to prevent early exit of ctas in a Cluster
CUTLASS_DEVICE void
load_tail(MainloopPipeline pipeline, MainloopPipelineState mainloop_pipe_producer_state) {
// Issue the epilogue waits
/* This helps avoid early exit of ctas in Cluster
* Waits for all stages to either be released (all
* Consumer UNLOCKs), or if the stage was never used
* then would just be acquired since the phase was
* still inverted from make_producer_start_state
*/
pipeline.producer_tail(mainloop_pipe_producer_state);
}
/// Perform a collective-scoped matrix multiply-accumulate
/// Consumer Perspective
template <
class FrgEngine, class FrgLayout,
class FragmentA, class FragmentB
>
CUTLASS_DEVICE auto
mma(MainloopPipeline pipeline,
MainloopPipelineState mainloop_pipe_consumer_state,
cute::Tensor<FrgEngine, FrgLayout>& accumulators,
cute::tuple<TiledMma, FragmentA, FragmentB> const& mma_inputs,
int k_tile_count)
{
static_assert(is_tmem<FrgEngine>::value, "Accumulator must be tmem resident.");
static_assert(rank(FrgLayout{}) == 3, "Accumulator must be MMA-partitioned: (MMA, MMA_M, MMA_N)");
auto [tiled_mma, tCrA, tCrB] = mma_inputs;
uint32_t skip_wait = k_tile_count <= 0;
auto barrier_token = pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
//
// PIPELINED MAIN LOOP
//
tiled_mma.accumulate_ = UMMA::ScaleOut::Zero;
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
// WAIT on mainloop_pipe_consumer_state until its data are available (phase bit flips from mainloop_pipe_consumer_state.phase() value)
pipeline.consumer_wait(mainloop_pipe_consumer_state, barrier_token);
// Compute on k_tile
int read_stage = mainloop_pipe_consumer_state.index();
// Save current mainlop pipeline read state
auto curr_mainloop_pipe_consumer_state = mainloop_pipe_consumer_state;
// Advance mainloop_pipe
++mainloop_pipe_consumer_state;
--k_tile_count;
skip_wait = k_tile_count <= 0;
// Peek at next iteration
barrier_token = pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
// Unroll the K mode manually so we can set scale C to 1
CUTLASS_PRAGMA_UNROLL
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
// (V,M,K) x (V,N,K) => (V,M,N)
cute::gemm(tiled_mma, tCrA(_,_,k_block,read_stage), tCrB(_,_,k_block,read_stage), accumulators);
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
}
pipeline.consumer_release(curr_mainloop_pipe_consumer_state);
}
return mainloop_pipe_consumer_state;
}
CUTLASS_DEVICE auto
mma_init(TensorStorage& shared_tensors) {
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.data()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.data()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
TiledMma tiled_mma;
// Allocate "fragments/descriptors" for A and B matrices
Tensor tCrA = tiled_mma.make_fragment_A(sA); // (MMA,MMA_M,MMA_K,PIPE)
Tensor tCrB = tiled_mma.make_fragment_B(sB); // (MMA,MMA_N,MMA_K,PIPE)
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sA)); // PIPE
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sB)); // PIPE
return cute::make_tuple(tiled_mma, tCrA, tCrB);
}
private:
typename Params::TMA_A const* observed_tma_load_a_ = nullptr;
typename Params::TMA_B const* observed_tma_load_b_ = nullptr;
};
/////////////////////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::conv::collective
/////////////////////////////////////////////////////////////////////////////////////////////////
@@ -326,7 +326,9 @@ public:
else {
CUTLASS_ASSERT(cuda_adapter == nullptr);
void const* kernel = (void const*) device_kernel<ConvKernel>;
if constexpr (ConvKernel::ArchTag::kMinComputeCapability == 90) {
if constexpr (ConvKernel::ArchTag::kMinComputeCapability == 90
|| ConvKernel::ArchTag::kMinComputeCapability == 100
) {
if constexpr (is_static_1x1x1) {
device_kernel<ConvKernel><<<grid, block, smem_size, stream>>>(params);
launch_result = Status::kSuccess;
+31
View File
@@ -83,6 +83,37 @@ struct MainloopSm90TmaGmmaWarpSpecializedImplicitGemm {
"Persistent schedules not support for conv yet.");
};
// SM100 tensor op kernel schedule
struct KernelImplicitTmaWarpSpecializedSm100 { };
// Pseudo-policies for builder auto override that dispatches to the KernelImplicitTmaWarpSpecializedSm100
// but for opting into 1 or 2 SM atoms
struct KernelImplicitTmaWarpSpecialized1SmSm100 : KernelImplicitTmaWarpSpecializedSm100 { };
struct KernelImplicitTmaWarpSpecialized2SmSm100 : KernelImplicitTmaWarpSpecializedSm100 { };
struct KernelStridedDgradTmaWs1SmSm100 { };
struct KernelStridedDgradTmaWs2SmSm100 { };
// n-buffer in smem (Blackwell TMA), pipelined with Blackwell UMMA and TMA, fprop
template<
conv::Operator ConvOp_,
int Stages_,
int NumSpatialDimensions_,
class ClusterShape_ = cute::Shape<cute::C<1>,cute::C<1>,cute::C<1>>
>
struct MainloopSm100TmaUmmaWarpSpecializedImplicitGemm {
static constexpr int Stages = Stages_;
static constexpr int NumSpatialDimensions = NumSpatialDimensions_;
static constexpr Operator ConvOp = ConvOp_;
using ClusterShape = ClusterShape_;
using ArchTag = arch::Sm100;
using Schedule = KernelImplicitTmaWarpSpecializedSm100;
static_assert(NumSpatialDimensions >= 1);
};
//////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::conv
@@ -61,4 +61,5 @@ class ConvUniversal {
////////////////////////////////////////////////////////////////////////////////
#include "cutlass/conv/kernel/sm90_implicit_gemm_tma_warpspecialized.hpp"
#include "cutlass/conv/kernel/sm100_implicit_gemm_tma_warpspecialized.hpp"
////////////////////////////////////////////////////////////////////////////////
@@ -0,0 +1,182 @@
/***************************************************************************************************
* Copyright (c) 2024 - 2025 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.
*
**************************************************************************************************/
#pragma once
#include "cutlass/conv/kernel/conv_universal.hpp"
#include "cutlass/gemm/kernel/tile_scheduler.hpp"
#include "cutlass/fast_math.h"
#include "cutlass/workspace.h"
#include <cute/util/type_traits.hpp>
#include <cute/int_tuple.hpp>
////////////////////////////////////////////////////////////////////////////////
namespace cutlass::conv::kernel {
////////////////////////////////////////////////////////////////////////////////
enum class DispatchMode {
VoidC // Select between voidC and non-voidC kernel based on beta scaling
};
// Dispatch between two ConvUniversal kernels
template <DispatchMode Mode, class KernelA, class KernelB, class = void>
class ConvUniversalDispatch;
////////////////////////////////////////////////////////////////////////////////
template <
class ProblemShape_,
class MainloopWithC_, class EpilogueWithC_,
class MainloopVoidC_, class EpilogueVoidC_,
class TileScheduler_
>
class ConvUniversalDispatch<
DispatchMode::VoidC,
ConvUniversal<ProblemShape_, MainloopWithC_, EpilogueWithC_, TileScheduler_>,
ConvUniversal<ProblemShape_, MainloopVoidC_, EpilogueVoidC_, TileScheduler_>,
cute::void_t<decltype(typename EpilogueWithC_::Arguments{}.thread.dBeta),
decltype(typename EpilogueVoidC_::Arguments{}.thread.dBeta)>
> : public ConvUniversal<ProblemShape_, MainloopWithC_, EpilogueWithC_, TileScheduler_> {
private:
using KernelWithC = ConvUniversal<ProblemShape_, MainloopWithC_, EpilogueWithC_, TileScheduler_>;
using KernelVoidC = ConvUniversal<ProblemShape_, MainloopVoidC_, EpilogueVoidC_, TileScheduler_>;
using FusionArguments = cute::remove_cvref_t<decltype(typename EpilogueWithC_::Arguments{}.thread)>;
public:
// Mainloop derived types
static_assert(cute::is_same_v<typename KernelWithC::TileShape, typename KernelVoidC::TileShape>);
static_assert(cute::is_same_v<typename KernelWithC::TiledMma, typename KernelVoidC::TiledMma>);
static_assert(cute::is_same_v<typename KernelWithC::ArchTag, typename KernelVoidC::ArchTag>);
static_assert(cute::is_same_v<typename KernelWithC::ElementA, typename KernelVoidC::ElementA>);
static_assert(cute::is_same_v<typename KernelWithC::StrideA, typename KernelVoidC::StrideA>);
static_assert(cute::is_same_v<typename KernelWithC::ElementB, typename KernelVoidC::ElementB>);
static_assert(cute::is_same_v<typename KernelWithC::StrideB, typename KernelVoidC::StrideB>);
static_assert(cute::is_same_v<typename KernelWithC::ElementAccumulator, typename KernelVoidC::ElementAccumulator>);
static_assert(cute::is_same_v<typename KernelWithC::ClusterShape, typename KernelVoidC::ClusterShape>);
// Epilogue derived types
static_assert(not cute::is_void_v<typename KernelWithC::ElementC>);
static_assert( cute::is_void_v<typename KernelVoidC::ElementC>);
static_assert(cute::is_same_v<typename KernelWithC::StrideC, typename KernelVoidC::StrideC>);
static_assert(cute::is_same_v<typename KernelWithC::ElementD, typename KernelVoidC::ElementD>);
static_assert(cute::is_same_v<typename KernelWithC::StrideD, typename KernelVoidC::StrideD>);
// TileID scheduler
static_assert(cute::is_same_v<typename KernelWithC::TileScheduler, typename KernelVoidC::TileScheduler>);
static constexpr int SharedStorageSize = cute::max(KernelWithC::SharedStorageSize, KernelVoidC::SharedStorageSize);
static_assert(KernelWithC::MaxThreadsPerBlock == KernelVoidC::MaxThreadsPerBlock);
static_assert(KernelWithC::MinBlocksPerMultiprocessor == KernelVoidC::MinBlocksPerMultiprocessor);
using Arguments = typename KernelWithC::Arguments;
struct Params {
typename KernelWithC::Params withC;
typename KernelVoidC::Params voidC;
void const* ptr_C;
decltype(FusionArguments{}.beta) beta;
decltype(FusionArguments{}.beta_ptr) beta_ptr;
decltype(FusionArguments{}.dBeta) dBeta;
cutlass::KernelHardwareInfo hw_info{};
};
static size_t
get_workspace_size(Arguments const& args) {
return KernelWithC::get_workspace_size(args);
}
static cutlass::Status
initialize_workspace(Arguments const& args, void* workspace = nullptr, cudaStream_t stream = nullptr, CudaHostAdapter* cuda_adapter = nullptr) {
return KernelWithC::initialize_workspace(args, workspace, stream, cuda_adapter);
}
static Params
to_underlying_arguments(Arguments const& args, void* workspace) {
return {
KernelWithC::to_underlying_arguments(args, workspace),
KernelVoidC::to_underlying_arguments(reinterpret_cast<typename KernelVoidC::Arguments const&>(args), workspace),
args.epilogue.ptr_C,
args.epilogue.thread.beta,
args.epilogue.thread.beta_ptr,
args.epilogue.thread.dBeta,
args.hw_info
};
}
static dim3
get_grid_shape(Params const& params) {
return KernelWithC::get_grid_shape(params.withC);
}
CUTLASS_DEVICE
void
operator()(Params const& params, char* smem_buf) {
using namespace cute;
bool run_voidC = false;
if (params.ptr_C == nullptr) {
run_voidC = true;
}
else if (params.beta_ptr == nullptr) { // Host scalar beta
run_voidC = params.beta == 0;
}
else if (get<0>(params.dBeta) == 0 && get<1>(params.dBeta) == 0) { // Device scalar beta
auto L = get<3>(append<4>(params.withC.problem_shape, _1{}));
if (get<2>(params.dBeta) == repeat_like(L, 0) || size(L) == 1) { // Non-batched
run_voidC = *params.beta_ptr == 0;
}
}
if (run_voidC) {
return kernel_voidC(params.voidC, smem_buf);
}
else {
return KernelWithC::operator()(params.withC, smem_buf);
}
}
private:
KernelVoidC kernel_voidC;
};
////////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::conv::kernel
////////////////////////////////////////////////////////////////////////////////
@@ -0,0 +1,911 @@
/***************************************************************************************************
* Copyright (c) 2023 - 2025 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.
*
**************************************************************************************************/
#pragma once
#include "cutlass/cutlass.h"
#include "cutlass/fast_math.h"
#include "cutlass/kernel_hardware_info.hpp"
#include "cute/tensor.hpp"
#include "cute/arch/tmem_allocator_sm100.hpp"
#include "cute/arch/cluster_sm90.hpp"
#include "cutlass/arch/arch.h"
#include "cutlass/arch/grid_dependency_control.h"
#include "cutlass/conv/detail.hpp"
#include "cutlass/conv/convolution.h"
#include "cutlass/conv/dispatch_policy.hpp"
#include "cutlass/gemm/kernel/tile_scheduler.hpp"
#include "cutlass/pipeline/sm100_pipeline.hpp"
#include "cutlass/detail/sm100_tmem_helper.hpp"
///////////////////////////////////////////////////////////////////////////////
namespace cutlass::conv::kernel {
///////////////////////////////////////////////////////////////////////////////
template <
class ProblemShape_,
class CollectiveMainloop_,
class CollectiveEpilogue_,
class TileSchedulerTag_
>
class ConvUniversal<
ProblemShape_,
CollectiveMainloop_,
CollectiveEpilogue_,
TileSchedulerTag_,
cute::enable_if_t<cute::is_base_of_v<KernelImplicitTmaWarpSpecializedSm100,
typename CollectiveMainloop_::DispatchPolicy::Schedule>>>
{
public:
//
// Type Aliases
//
// Mainloop derived types
using ProblemShape = ProblemShape_;
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 DispatchPolicy = typename CollectiveMainloop::DispatchPolicy;
using ElementAccumulator = typename CollectiveMainloop::ElementAccumulator;
using ClusterShape = typename DispatchPolicy::ClusterShape;
using MainloopArguments = typename CollectiveMainloop::Arguments;
using MainloopParams = typename CollectiveMainloop::Params;
using CtaShape_MNK = typename CollectiveMainloop::CtaShape_MNK;
using AtomThrShapeMNK = typename CollectiveMainloop::AtomThrShapeMNK;
static constexpr int NumSpatialDimensions = CollectiveMainloop::NumSpatialDimensions;
static constexpr bool is_grouped_wgrad = CollectiveMainloop::is_grouped_wgrad;
static constexpr bool IsComplex = false;
static_assert(ArchTag::kMinComputeCapability >= 100);
// Epilogue derived types
using CollectiveEpilogue = CollectiveEpilogue_;
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 IsGdcEnabled = cutlass::arch::IsGdcGloballyEnabled;
// TileID scheduler
// CLC pipeline depth determines how many waves (stages-1) the scheduler can race ahead
static constexpr uint32_t SchedulerPipelineStageCount = 2;
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 IsDynamicCluster = not cute::is_static_v<ClusterShape>;
// Warp specialization thread count per threadblock
static constexpr uint32_t NumSchedThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumMMAThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumMainloopLoadThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumEpilogueLoadThreads = NumThreadsPerWarp; // 1 warp
static constexpr uint32_t NumEpilogueThreads = CollectiveEpilogue::ThreadCount;
static constexpr uint32_t NumEpilogueWarps = NumEpilogueThreads / NumThreadsPerWarp;
static constexpr uint32_t MaxThreadsPerBlock = NumSchedThreads +
NumMainloopLoadThreads + NumMMAThreads +
NumEpilogueLoadThreads + NumEpilogueThreads;
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
static constexpr uint32_t NumFixupBarriers = 1;
// Pipelines and pipeline states
static constexpr uint32_t AccumulatorPipelineStageCount = SchedulerPipelineStageCount;
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 = cutlass::PipelineDetail::PipelineCLCFetchAsyncPipelineState<SchedulerPipelineStageCount>;
using CLCPipelineSharedStorage = cutlass::PipelineDetail::PipelineCLCFetchAsyncSharedStorage<SchedulerPipelineStageCount>;
using CLCThrottlePipeline = cutlass::PipelineAsync<SchedulerPipelineStageCount>;
using CLCThrottlePipelineState = cutlass::PipelineDetail::PipelineAsyncPipelineState<SchedulerPipelineStageCount>;
using CLCThrottlePipelineSharedStorage = cutlass::PipelineDetail::PipelineAsyncSharedStorage<SchedulerPipelineStageCount>;
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 = CLCPipelineSharedStorage;
using AccumulatorPipelineStorage = typename AccumulatorPipeline::SharedStorage;
using CLCThrottlePipelineStorage = CLCThrottlePipelineSharedStorage;
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) 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);
static_assert(SharedStorageSize <= cutlass::arch::sm100_smem_capacity_bytes, "SMEM usage exceeded capacity.");
// Host facing host arguments
struct Arguments {
ProblemShape problem_shape{};
MainloopArguments mainloop{};
EpilogueArguments epilogue{};
KernelHardwareInfo hw_info{};
TileSchedulerArguments scheduler{};
};
// Kernel device entry point API
struct Params {
using ProblemShapeMNKL = decltype(CollectiveMainloop::get_problem_shape_MNKL(ProblemShape{}));
ProblemShapeMNKL problem_shape;
MainloopParams mainloop;
EpilogueParams epilogue;
TileSchedulerParams scheduler;
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
//
// Map user facing arguments to device facing params
CUTLASS_HOST
static Params
to_underlying_arguments(Arguments const& args, void* workspace) {
static constexpr uint32_t NumEpilogueSubTiles = 1;
auto problem_shape_mnkl = CollectiveMainloop::get_problem_shape_MNKL(args.problem_shape);
auto mainloop_params = CollectiveMainloop::to_underlying_arguments(args.problem_shape, args.mainloop, workspace, args.hw_info);
// 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 = round_nearest(workspace_offset, MinWorkspaceAlignment);
// Tile scheduler
void* scheduler_workspace = workspace_ptr + workspace_offset;
workspace_offset += TileScheduler::template get_workspace_size<decltype(problem_shape_mnkl), ElementAccumulator>(
args.scheduler, problem_shape_mnkl, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
return {
problem_shape_mnkl,
mainloop_params,
CollectiveEpilogue::to_underlying_arguments(args.problem_shape, args.epilogue, epilogue_workspace),
TileScheduler::to_underlying_arguments(
args.problem_shape, TileShape{}, AtomThrShapeMNK{}, ClusterShape{},
args.hw_info, args.scheduler, scheduler_workspace),
args.hw_info
};
}
CUTLASS_HOST
static bool
can_implement(Arguments const& args) {
bool implementable = true;
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);
}
if constexpr (is_grouped_wgrad) {
auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, args.hw_info.cluster_shape);
auto cluster_shape_fallback = cutlass::detail::select_cluster_shape(ClusterShape{}, args.hw_info.cluster_shape_fallback);
implementable &= size<0>(cluster_shape) == 1 && size<0>(cluster_shape_fallback) == 1;
if (!implementable) {
return false;
}
}
return implementable;
}
CUTLASS_HOST
static size_t
get_workspace_size(Arguments const& args) {
static constexpr uint32_t NumEpilogueSubTiles = 1;
size_t workspace_size = 0;
auto linear_problem_shape_MNKL = cutlass::conv::detail::get_linearized_problem_shape_MNKL(args.problem_shape);
// Epilogue
workspace_size += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
// Tile scheduler
workspace_size += TileScheduler::template get_workspace_size<decltype(linear_problem_shape_MNKL), ElementAccumulator>(
args.scheduler, linear_problem_shape_MNKL, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs);
workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment);
return workspace_size;
}
CUTLASS_HOST
static cutlass::Status
initialize_workspace(Arguments const& args, void* workspace = nullptr, cudaStream_t stream = nullptr,
CudaHostAdapter* cuda_adapter = nullptr) {
static constexpr uint32_t NumEpilogueSubTiles = 1;
auto linear_problem_shape_MNKL = cutlass::conv::detail::get_linearized_problem_shape_MNKL(args.problem_shape);
Status status = 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, stream, cuda_adapter);
workspace_offset += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
// Tile scheduler
status = TileScheduler::template initialize_workspace
<decltype(linear_problem_shape_MNKL), ElementAccumulator>(
args.scheduler, workspace_ptr + workspace_offset, stream, linear_problem_shape_MNKL,
args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs, cuda_adapter);
workspace_offset += TileScheduler::template get_workspace_size
<decltype(linear_problem_shape_MNKL), ElementAccumulator>(
args.scheduler, linear_problem_shape_MNKL, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles,
CollectiveEpilogue::NumAccumulatorMtxs);
workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment);
if (status != Status::kSuccess) {
return status;
}
return status;
}
// Computes the kernel launch grid shape based on runtime parameters
CUTLASS_HOST
static dim3
get_grid_shape(Params const& params) {
auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, params.hw_info.cluster_shape);
return TileScheduler::get_grid_shape(
params.scheduler,
params.problem_shape,
TileShape{},
AtomThrShapeMNK{},
cluster_shape
,params.hw_info
);
}
CUTLASS_HOST
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;
// Separate out problem shape for convenience
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 = 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{}, cute::cluster_shape());
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;
// Issue Tma Descriptor Prefetch from a single thread
if ((warp_category == WarpCategory::Sched) && lane_predicate) {
CollectiveMainloop::prefetch_tma_descriptors(params.mainloop);
}
if ((warp_category == WarpCategory::EpilogueLoad) && lane_predicate) {
CollectiveEpilogue::prefetch_tma_descriptors(params.epilogue);
}
// 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);
CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.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 = 4;
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 = 5;
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 = 1;
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 = 2;
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
arch::NamedBarrier tmem_allocation_result_barrier(NumMMAThreads + NumEpilogueThreads, cutlass::arch::ReservedNamedBarriers::TmemAllocBarrier);
// Sync deallocation status between MMA warps of peer CTAs
arch::ClusterBarrier& tmem_deallocation_result_barrier = shared_storage.pipelines.tmem_dealloc;
[[maybe_unused]] uint32_t dealloc_barrier_phase = 0;
if (WarpCategory::MMA == warp_category && has_mma_peer_cta && lane_predicate) {
tmem_deallocation_result_barrier.init(NumMMAThreads);
}
// We need this to guarantee that the Pipeline init is visible
// To all producers and consumer threadblocks in the cluster
if (cluster_size > 1) {
cute::cluster_arrive_relaxed();
}
else {
__syncthreads();
}
uint32_t tmem_stage_ptrs[AccumulatorPipelineStageCount];
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);
// TileID scheduler
TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, problem_shape_MNKL, TileShape{}, 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);
auto acc_shape = collective_mainloop.partition_accumulator_shape();
auto accumulators = TiledMma::make_fragment_C(acc_shape);
int TmemColumnsPerAccumulatorTile = cutlass::detail::find_tmem_tensor_col_offset(accumulators);
pipeline_init_wait(cluster_size);
if (is_participant.sched) {
// 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;
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 stream-K units.
// The current stream-K scheduler uses single WorkTileInfo to track multiple
// (potentially-partial) tiles to be computed via stream-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.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;
auto load_inputs = collective_mainloop.load_init(
problem_shape_MNKL, params.mainloop, shared_storage.tensors.mainloop);
Tensor gA_mk = get<0>(load_inputs);
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, TileShape{}, shape<3>(gA_mk));
auto k_tile_count = scheduler.get_work_k_tile_count(work_tile_info, problem_shape_MNKL, TileShape{});
auto k_tile_prologue = min(MainloopPipeline::Stages, k_tile_count);
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;
}
auto [mainloop_producer_state_next, k_tile_iter_next] = collective_mainloop.load(
params.mainloop,
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(
params.mainloop,
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.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;
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;
}
epi_load_pipe_producer_state = collective_epilogue.load(
epi_load_pipeline,
epi_load_pipe_producer_state,
problem_shape_MNKL,
CtaShape_MNK{},
cta_coord_mnkl,
TileShape{},
TiledMma{},
shared_storage.tensors.epilogue
);
do_tail_load = true;
}
// 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());
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.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;
CUTLASS_PRAGMA_UNROLL
for (int acc_stage = 0; acc_stage < AccumulatorPipelineStageCount; acc_stage++) {
tmem_stage_ptrs[acc_stage] = tmem_base_ptr + (TmemColumnsPerAccumulatorTile * acc_stage) & cutlass::detail::TmemColMask;
}
auto mma_inputs = collective_mainloop.mma_init(shared_storage.tensors.mainloop);
do {
auto k_tile_count = scheduler.get_work_k_tile_count(work_tile_info, problem_shape_MNKL, TileShape{});
// 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;
}
// Wait for tmem accumulator buffer to become empty with a flipped phase
if (is_mma_leader_cta) {
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
}
// Accumulator stage slice
int acc_stage = accumulator_pipe_producer_state.index();
accumulators.data() = tmem_stage_ptrs[acc_stage];
if (is_mma_leader_cta) {
mainloop_pipe_consumer_state = collective_mainloop.mma(
mainloop_pipeline,
mainloop_pipe_consumer_state,
accumulators,
mma_inputs,
k_tile_count
);
accumulator_pipeline.producer_commit(accumulator_pipe_producer_state);
}
++accumulator_pipe_producer_state;
work_tile_info = next_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();
// 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);
}
// Free entire tmem allocation
tmem_allocator.free(tmem_base_ptr, TmemAllocator::Sm100TmemCapacityColumns);
}
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;
CUTLASS_PRAGMA_UNROLL
for (int acc_stage = 0; acc_stage < AccumulatorPipelineStageCount; acc_stage++) {
tmem_stage_ptrs[acc_stage] = tmem_base_ptr + (TmemColumnsPerAccumulatorTile * acc_stage) & cutlass::detail::TmemColMask;
}
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 after making sure allocation has been performed
int acc_stage = accumulator_pipe_consumer_state.index();
accumulators.data() = tmem_stage_ptrs[acc_stage];
accumulator_pipe_consumer_state = scheduler.template fixup<IsComplex>(
TiledMma{},
work_tile_info,
accumulators,
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.store(
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{},
accumulators,
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 (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{});
}
}
}
private:
// Synchronization call. Blocks until barriers are initialized in shared memory.
CUTLASS_DEVICE
void
pipeline_init_wait(int cluster_size) {
if (cluster_size > 1) {
cute::cluster_wait();
}
else {
__syncthreads();
}
}
};
///////////////////////////////////////////////////////////////////////////////
} // namespace cutlass::gemm::kernel