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:
co-authored by
Haicheng Wu
Haicheng Wu
parent
9eb01fa0b0
commit
389e493055
@@ -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"
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -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 ? ¶ms.tma_load_a_fallback : ¶ms.tma_load_a;
|
||||
observed_tma_load_b_ = is_fallback_cluster ? ¶ms.tma_load_b_fallback : ¶ms.tma_load_b;
|
||||
}
|
||||
else {
|
||||
observed_tma_load_a_ = ¶ms.tma_load_a;
|
||||
observed_tma_load_b_ = ¶ms.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;
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user