CUTLASS 3.5.0 (#1411)
This commit is contained in:
@@ -0,0 +1,96 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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 "cutlass/arch/mma.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/dispatch_policy.hpp"
|
||||
#include "cutlass/detail/layout.hpp"
|
||||
#include "cutlass/gemm/collective/builders/sm90_common.inl"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv::collective::detail {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Maps a rank-1 cute::Shape<> representing the cluster shape on to the IM2COL TMA atom that should be used with it
|
||||
template <class UnimodalClusterShape>
|
||||
constexpr auto
|
||||
sm90_cluster_shape_to_im2col_tma_atom(UnimodalClusterShape unimodal_cluster_shape) {
|
||||
static_assert(cute::rank(unimodal_cluster_shape) == 1,
|
||||
"Use this function to figure out TMA for each mode individually.");
|
||||
|
||||
if constexpr (cute::size(unimodal_cluster_shape) == 1) {
|
||||
return cute::SM90_TMA_LOAD_IM2COL{};
|
||||
}
|
||||
else {
|
||||
return cute::SM90_TMA_LOAD_IM2COL_MULTICAST{};
|
||||
}
|
||||
}
|
||||
|
||||
// Collective tile traits struct that serves as a type list containing a tensor's mem layouts and atoms for the
|
||||
template<
|
||||
class GmemTiledCopy_,
|
||||
class SmemLayout_,
|
||||
class SmemCopyAtom_ = void
|
||||
>
|
||||
struct Sm90ImplicitGemmTileTraits {
|
||||
using GmemTiledCopy = GmemTiledCopy_;
|
||||
using SmemLayout = SmemLayout_;
|
||||
using SmemCopyAtom = SmemCopyAtom_;
|
||||
};
|
||||
|
||||
// Accepts a cutlass::layout::Tensor tag and computes the corresponding spatial dimension count
|
||||
template <class GmemLayoutTagA, class GmemLayoutTagB>
|
||||
constexpr int
|
||||
gmem_layout_tags_to_spatial_dims() {
|
||||
static_assert(cute::is_same_v<GmemLayoutTagA, GmemLayoutTagB>);
|
||||
if constexpr (cute::is_same_v<GmemLayoutTagA, cutlass::layout::TensorNWC>) {
|
||||
return 1;
|
||||
}
|
||||
else if constexpr (cute::is_same_v<GmemLayoutTagA, cutlass::layout::TensorNHWC>) {
|
||||
return 2;
|
||||
}
|
||||
else if constexpr (cute::is_same_v<GmemLayoutTagA, cutlass::layout::TensorNDHWC>) {
|
||||
return 3;
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<GmemLayoutTagA>);
|
||||
}
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::collective::detail
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,257 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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/sm90_common.inl"
|
||||
|
||||
// SM90 Collective Builders should be used only starting CUDA 12.0
|
||||
#if (__CUDACC_VER_MAJOR__ >= 12)
|
||||
#define CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
#endif
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv::collective {
|
||||
using namespace cute;
|
||||
|
||||
namespace detail {
|
||||
|
||||
// Returns the maximum number of smem tiles that can be used with a given smem capacity, or overrides with manual count.
|
||||
template<int CapacityBytes, class ElementA, class ElementB, class TileShapeMNK, int stages>
|
||||
constexpr int
|
||||
compute_stage_count_or_override(StageCount<stages> stage_count) {
|
||||
return stages;
|
||||
}
|
||||
|
||||
// Returns the maximum number of smem tiles that can be used with a given smem capacity, or overrides with manual count.
|
||||
template<int CapacityBytes, class ElementA, class ElementB, class TileShapeMNK, int stages>
|
||||
constexpr int
|
||||
compute_stage_count_or_override(cute::Int<stages> stage_count) {
|
||||
return stages;
|
||||
}
|
||||
|
||||
// Returns the maximum number of smem tiles that can be used with a given smem capacity, or overrides with manual count.
|
||||
template<int CapacityBytes, class ElementA, class ElementB, class TileShapeMNK, int carveout_bytes>
|
||||
constexpr int
|
||||
compute_stage_count_or_override(StageCountAutoCarveout<carveout_bytes> stage_count) {
|
||||
constexpr auto mainloop_pipeline_bytes = sizeof(typename cutlass::PipelineTmaAsync<1>::SharedStorage);
|
||||
constexpr auto a_bits = cute::sizeof_bits_v<ElementA>;
|
||||
constexpr auto b_bits = cute::sizeof_bits_v<ElementB>;
|
||||
constexpr int stage_bytes =
|
||||
cutlass::bits_to_bytes(a_bits * size<0>(TileShapeMNK{}) * size<2>(TileShapeMNK{})) +
|
||||
cutlass::bits_to_bytes(b_bits * size<1>(TileShapeMNK{}) * size<2>(TileShapeMNK{})) +
|
||||
static_cast<int>(mainloop_pipeline_bytes);
|
||||
|
||||
return (CapacityBytes - carveout_bytes) / stage_bytes;
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// GMMA_TMA_WS_SS_FPROP
|
||||
template <
|
||||
conv::Operator ConvOp,
|
||||
class ElementA,
|
||||
class GmemLayoutA,
|
||||
int AlignmentA,
|
||||
class ElementB,
|
||||
class GmemLayoutB,
|
||||
int AlignmentB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
class StageCountType,
|
||||
class KernelScheduleType
|
||||
>
|
||||
struct CollectiveBuilder<
|
||||
arch::Sm90,
|
||||
arch::OpClassTensorOp,
|
||||
ConvOp,
|
||||
ElementA,
|
||||
GmemLayoutA,
|
||||
AlignmentA,
|
||||
ElementB,
|
||||
GmemLayoutB,
|
||||
AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape_MNK,
|
||||
ClusterShape_MNK,
|
||||
StageCountType,
|
||||
KernelScheduleType,
|
||||
cute::enable_if_t<cute::is_same_v<KernelScheduleType, KernelImplicitTmaWarpSpecializedSm90> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelImplicitTmaWarpSpecializedSm90Cooperative> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelImplicitTmaWarpSpecializedSm90Pingpong>>
|
||||
> {
|
||||
static_assert(is_static<TileShape_MNK>::value);
|
||||
static_assert(is_static<ClusterShape_MNK>::value);
|
||||
#ifndef CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
#endif
|
||||
static_assert(cutlass::gemm::collective::detail::is_aligned<ElementA, AlignmentA, ElementB, AlignmentB, cutlass::gemm::collective::detail::tma_alignment_bytes>(),
|
||||
"Should meet TMA alignment requirement\n");
|
||||
|
||||
// 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>;
|
||||
|
||||
// For fprop, majorA = K, major B = K;
|
||||
// For wgrad, majorA = MN, major B = MN;
|
||||
// For dgrad, majorA = K, major B = MN;
|
||||
static constexpr cute::GMMA::Major GmmaMajorA =
|
||||
(ConvOp == conv::Operator::kWgrad) ? cute::GMMA::Major::MN : cute::GMMA::Major::K;
|
||||
static constexpr cute::GMMA::Major GmmaMajorB =
|
||||
(ConvOp == conv::Operator::kFprop) ? cute::GMMA::Major::K : cute::GMMA::Major::MN;
|
||||
|
||||
using AtomLayoutMNK = cute::conditional_t<cute::is_same_v<KernelScheduleType, KernelImplicitTmaWarpSpecializedSm90Cooperative>,
|
||||
Layout<Shape<_2,_1,_1>>, Layout<Shape<_1,_1,_1>>>;
|
||||
|
||||
using TiledMma = decltype(cute::make_tiled_mma(cute::GMMA::ss_op_selector<
|
||||
ElementAMma, ElementBMma, ElementAccumulator, TileShape_MNK, GmmaMajorA, GmmaMajorB>(), AtomLayoutMNK{}));
|
||||
|
||||
// For wgrad kernel, tensor A uses tma tiled mode and tensor B uses tma im2col mode.
|
||||
using GmemTiledCopyA = cute::conditional_t<ConvOp == conv::Operator::kWgrad,
|
||||
decltype(cutlass::gemm::collective::detail::sm90_cluster_shape_to_tma_atom(cute::shape<1>(ClusterShape_MNK{}))),
|
||||
decltype(cutlass::conv::collective::detail::sm90_cluster_shape_to_im2col_tma_atom(cute::shape<1>(ClusterShape_MNK{})))>;
|
||||
using GmemTiledCopyB = cute::conditional_t<ConvOp == conv::Operator::kWgrad,
|
||||
decltype(cutlass::conv::collective::detail::sm90_cluster_shape_to_im2col_tma_atom(cute::shape<0>(ClusterShape_MNK{}))),
|
||||
decltype(cutlass::gemm::collective::detail::sm90_cluster_shape_to_tma_atom(cute::shape<0>(ClusterShape_MNK{})))>;
|
||||
|
||||
using SmemLayoutAtomA = decltype(cutlass::gemm::collective::detail::ss_smem_selector<
|
||||
GmmaMajorA, ElementAMma, decltype(cute::get<0>(TileShape_MNK{})), decltype(cute::get<2>(TileShape_MNK{}))>());
|
||||
using SmemLayoutAtomB = decltype(cutlass::gemm::collective::detail::ss_smem_selector<
|
||||
GmmaMajorB, ElementBMma, decltype(cute::get<1>(TileShape_MNK{})), decltype(cute::get<2>(TileShape_MNK{}))>());
|
||||
|
||||
static constexpr int PipelineStages = detail::compute_stage_count_or_override<cutlass::gemm::collective::detail::sm90_smem_capacity_bytes,
|
||||
ElementAMma, ElementBMma, TileShape_MNK>(StageCountType{});
|
||||
|
||||
using SmemLayoutA = decltype(tile_to_shape(
|
||||
SmemLayoutAtomA{},
|
||||
make_shape(shape<0>(TileShape_MNK{}), shape<2>(TileShape_MNK{}), Int<PipelineStages>{}),
|
||||
Step<_2,_1,_3>{}));
|
||||
using SmemLayoutB = decltype(tile_to_shape(
|
||||
SmemLayoutAtomB{},
|
||||
make_shape(shape<1>(TileShape_MNK{}), shape<2>(TileShape_MNK{}), Int<PipelineStages>{}),
|
||||
Step<_2,_1,_3>{}));
|
||||
|
||||
constexpr static int NumSpatialDimensions = cutlass::conv::collective::detail::gmem_layout_tags_to_spatial_dims<GmemLayoutA, GmemLayoutB>();
|
||||
|
||||
using DispatchPolicy = MainloopSm90TmaGmmaWarpSpecializedImplicitGemm<
|
||||
ConvOp, PipelineStages, NumSpatialDimensions, ClusterShape_MNK, KernelScheduleType>;
|
||||
|
||||
using CollectiveOp = CollectiveConv<
|
||||
DispatchPolicy,
|
||||
TileShape_MNK,
|
||||
ElementA,
|
||||
ElementB,
|
||||
TiledMma,
|
||||
detail::Sm90ImplicitGemmTileTraits<GmemTiledCopyA, SmemLayoutA>,
|
||||
detail::Sm90ImplicitGemmTileTraits<GmemTiledCopyB, SmemLayoutB>
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// GMMA auto kernel schedule
|
||||
template <
|
||||
conv::Operator ConvOp,
|
||||
class ElementA,
|
||||
class GmemLayoutA,
|
||||
int AlignmentA,
|
||||
class ElementB,
|
||||
class GmemLayoutB,
|
||||
int AlignmentB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
class StageCountType,
|
||||
class KernelScheduleType
|
||||
>
|
||||
struct CollectiveBuilder<
|
||||
arch::Sm90,
|
||||
arch::OpClassTensorOp,
|
||||
ConvOp,
|
||||
ElementA,
|
||||
GmemLayoutA,
|
||||
AlignmentA,
|
||||
ElementB,
|
||||
GmemLayoutB,
|
||||
AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape_MNK,
|
||||
ClusterShape_MNK,
|
||||
StageCountType,
|
||||
KernelScheduleType,
|
||||
cute::enable_if_t<cute::is_same_v<KernelScheduleType, KernelScheduleAuto>>
|
||||
> {
|
||||
static_assert(is_static<TileShape_MNK>::value);
|
||||
static_assert(is_static<ClusterShape_MNK>::value);
|
||||
#ifndef CUTLASS_SM90_COLLECTIVE_BUILDER_SUPPORTED
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Unsupported Toolkit for SM90 Collective Builder\n");
|
||||
#endif
|
||||
|
||||
/*
|
||||
#if ((__CUDACC_VER_MAJOR__ > 12) || ((__CUDACC_VER_MAJOR__ == 12) && (__CUDACC_VER_MINOR__ >= 1)))
|
||||
// Cooperative schedule performs best for CUDA Toolkits with version >= 12.1
|
||||
|
||||
// For TileShape_M == 64, choosing KernelTmaWarpSpecialized as the KernelSchedule
|
||||
// Since KernelTmaWarpSpecializedCooperative requires TileShape_M to be at least 128
|
||||
using KernelWarpSpecializedSchedule = cute::conditional_t<size<0>(TileShape_MNK{}) == Int<64>{},
|
||||
KernelImplicitTmaWarpSpecializedSm90PingPong, KernelImplicitTmaWarpSpecializedSm90Cooperative>;
|
||||
#else
|
||||
using KernelWarpSpecializedSchedule = KernelImplicitTmaWarpSpecializedSm90;
|
||||
#endif
|
||||
*/
|
||||
using KernelWarpSpecializedSchedule = KernelImplicitTmaWarpSpecializedSm90;
|
||||
|
||||
using CollectiveOp = typename CollectiveBuilder<
|
||||
arch::Sm90,
|
||||
arch::OpClassTensorOp,
|
||||
ConvOp,
|
||||
ElementA,
|
||||
GmemLayoutA,
|
||||
AlignmentA,
|
||||
ElementB,
|
||||
GmemLayoutB,
|
||||
AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape_MNK,
|
||||
ClusterShape_MNK,
|
||||
StageCountType,
|
||||
KernelWarpSpecializedSchedule
|
||||
>::CollectiveOp;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,93 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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/detail/dependent_false.hpp"
|
||||
#include "cutlass/conv/collective/collective_conv.hpp"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv::collective {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Used to specify stage counts or dispatch to automatic computation of stage count
|
||||
template<int num_stages>
|
||||
struct StageCount {
|
||||
static constexpr int value = num_stages;
|
||||
|
||||
StageCount() = default;
|
||||
explicit StageCount(cute::Int<num_stages>) {}
|
||||
};
|
||||
|
||||
template<int carveout_bytes>
|
||||
struct StageCountAutoCarveout {
|
||||
static constexpr int bytes = carveout_bytes;
|
||||
|
||||
StageCountAutoCarveout() = default;
|
||||
explicit StageCountAutoCarveout(cute::Int<carveout_bytes>) {}
|
||||
};
|
||||
|
||||
// Used to automatically let the builder pick the kernel schedule.
|
||||
// Can be overridden with kernel schedule tags in cutlass/conv/dispatch_policy.hpp
|
||||
struct KernelScheduleAuto {};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
class ArchTag,
|
||||
class OpClass,
|
||||
conv::Operator,
|
||||
class ElementA,
|
||||
class GmemLayoutA,
|
||||
int AlignmentA,
|
||||
class ElementB,
|
||||
class GmemLayoutB,
|
||||
int AlignmentB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
class StageCountType,
|
||||
class KernelScheduleType,
|
||||
class Enable = void
|
||||
>
|
||||
struct CollectiveBuilder {
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Could not build a collective for given parameters.");
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include "builders/sm90_gmma_builder.inl"
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,62 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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/detail/dependent_false.hpp"
|
||||
#include "cutlass/conv/collective/detail.hpp"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv::collective {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
class DispatchPolicy,
|
||||
class TileShape,
|
||||
class ElementA,
|
||||
class ElementB,
|
||||
class TiledMma,
|
||||
class TileTraitsA,
|
||||
class TileTraitsB
|
||||
>
|
||||
struct CollectiveConv {
|
||||
static_assert(cutlass::detail::dependent_false<ElementA>, "Could not find a mainloop specialization.");
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include "sm90_implicit_gemm_gmma_ss_warpspecialized.hpp"
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,251 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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/convnd_problem_shape.hpp"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv::collective::detail {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Construct the stride types for conv collectives based on the dispatch policy, strides 64b by default
|
||||
template <class DispatchPolicy>
|
||||
constexpr auto
|
||||
sm90_dispatch_policy_to_stride_A() {
|
||||
if constexpr (DispatchPolicy::ConvOp == conv::Operator::kFprop) {
|
||||
// Maps to modes ((w,n), C)
|
||||
if constexpr (DispatchPolicy::NumSpatialDimensions == 1) {
|
||||
return cute::Stride<cute::Stride<int64_t, int64_t>,
|
||||
cute::Int<1>>{};
|
||||
}
|
||||
// Maps to modes ((w,h,n), C)
|
||||
else if constexpr (DispatchPolicy::NumSpatialDimensions == 2) {
|
||||
return cute::Stride<cute::Stride<int64_t, int64_t, int64_t>,
|
||||
cute::Int<1>>{};
|
||||
}
|
||||
// Maps to modes ((w,h,d,n), C)
|
||||
else if constexpr (DispatchPolicy::NumSpatialDimensions == 3) {
|
||||
return cute::Stride<cute::Stride<int64_t, int64_t, int64_t, int64_t>,
|
||||
cute::Int<1>>{};
|
||||
}
|
||||
// error dims assert
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<DispatchPolicy>, "Unsupported spatial dim count.");
|
||||
}
|
||||
}
|
||||
else if constexpr (DispatchPolicy::ConvOp == conv::Operator::kWgrad) {
|
||||
// Maps to modes (k, nq/npq/nzpq)
|
||||
if constexpr (DispatchPolicy::NumSpatialDimensions == 1 ||
|
||||
DispatchPolicy::NumSpatialDimensions == 2 ||
|
||||
DispatchPolicy::NumSpatialDimensions == 3) {
|
||||
return cute::Stride<cute::Int<1>, int64_t>{};
|
||||
}
|
||||
// error dims assert
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<DispatchPolicy>, "Unsupported spatial dim count.");
|
||||
}
|
||||
}
|
||||
else if constexpr (DispatchPolicy::ConvOp == conv::Operator::kDgrad) {
|
||||
// Maps to modes ((q,n), K)
|
||||
if constexpr (DispatchPolicy::NumSpatialDimensions == 1) {
|
||||
return cute::Stride<cute::Stride<int64_t, int64_t>,
|
||||
cute::Int<1>>{};
|
||||
}
|
||||
// Maps to modes ((q,p,n), K)
|
||||
else if constexpr (DispatchPolicy::NumSpatialDimensions == 2) {
|
||||
return cute::Stride<cute::Stride<int64_t, int64_t, int64_t>,
|
||||
cute::Int<1>>{};
|
||||
}
|
||||
// Maps to modes ((q,p,z,n), K)
|
||||
else if constexpr (DispatchPolicy::NumSpatialDimensions == 3) {
|
||||
return cute::Stride<cute::Stride<int64_t, int64_t, int64_t, int64_t>,
|
||||
cute::Int<1>>{};
|
||||
}
|
||||
// error dims assert
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<DispatchPolicy>, "Unsupported spatial dim count.");
|
||||
}
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<DispatchPolicy>, "Unsupported ConvOp.");
|
||||
}
|
||||
}
|
||||
|
||||
// Construct the stirde types for conv collectives based on the dispatch policy, strides 64b by default
|
||||
template <class DispatchPolicy>
|
||||
constexpr auto
|
||||
sm90_dispatch_policy_to_stride_B() {
|
||||
if constexpr (DispatchPolicy::ConvOp == conv::Operator::kFprop) {
|
||||
// Maps to modes (k, (C,s))
|
||||
if constexpr (DispatchPolicy::NumSpatialDimensions == 1) {
|
||||
return cute::Stride<int64_t, cute::Stride<cute::Int<1>, int64_t>>{};
|
||||
}
|
||||
// Maps to modes (k, (C,s,r))
|
||||
else if constexpr (DispatchPolicy::NumSpatialDimensions == 2) {
|
||||
return cute::Stride<int64_t, cute::Stride<cute::Int<1>, int64_t, int64_t>>{};
|
||||
}
|
||||
// Maps to modes (k, (C,s,r,t))
|
||||
else if constexpr (DispatchPolicy::NumSpatialDimensions == 3) {
|
||||
return cute::Stride<int64_t, cute::Stride<cute::Int<1>, int64_t, int64_t, int64_t>>{};
|
||||
}
|
||||
// error dims assert
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<DispatchPolicy>, "Unsupported spatial dim count.");
|
||||
}
|
||||
}
|
||||
else if constexpr (DispatchPolicy::ConvOp == conv::Operator::kWgrad) {
|
||||
// Maps to modes (C, (w,n))
|
||||
if constexpr (DispatchPolicy::NumSpatialDimensions == 1) {
|
||||
return cute::Stride<cute::Int<1>,
|
||||
cute::Stride<int64_t, int64_t>>{};
|
||||
}
|
||||
// Maps to modes (C, (w,h,n))
|
||||
else if constexpr (DispatchPolicy::NumSpatialDimensions == 2) {
|
||||
return cute::Stride<cute::Int<1>,
|
||||
cute::Stride<int64_t, int64_t, int64_t>>{};
|
||||
}
|
||||
// Maps to modes (C, (w,h,d,n))
|
||||
else if constexpr (DispatchPolicy::NumSpatialDimensions == 3) {
|
||||
return cute::Stride<cute::Int<1>,
|
||||
cute::Stride<int64_t, int64_t, int64_t, int64_t>>{};
|
||||
}
|
||||
// error dims assert
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<DispatchPolicy>, "Unsupported spatial dim count.");
|
||||
}
|
||||
}
|
||||
else if constexpr (DispatchPolicy::ConvOp == conv::Operator::kDgrad) {
|
||||
// Maps to modes (C, (k,s))
|
||||
if constexpr (DispatchPolicy::NumSpatialDimensions == 1) {
|
||||
return cute::Stride<cute::Int<1>, cute::Stride<int64_t, int64_t>>{};
|
||||
}
|
||||
// Maps to modes (C, (k,s,r))
|
||||
else if constexpr (DispatchPolicy::NumSpatialDimensions == 2) {
|
||||
return cute::Stride<cute::Int<1>, cute::Stride<int64_t, int64_t, int64_t>>{};
|
||||
}
|
||||
// Maps to modes (C, (k,s,r,t))
|
||||
else if constexpr (DispatchPolicy::NumSpatialDimensions == 3) {
|
||||
return cute::Stride<cute::Int<1>, cute::Stride<int64_t, int64_t, int64_t, int64_t>>{};
|
||||
}
|
||||
// error dims assert
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<DispatchPolicy>, "Unsupported spatial dim count.");
|
||||
}
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<DispatchPolicy>, "Unsupported ConvOp.");
|
||||
}
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Compute the lower/near corner, returning it as a cute::array in [W,H,D] order
|
||||
template <conv::Operator ConvOp, int NumSpatialDimensions>
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
compute_lower_corner_whd(ConvProblemShape<ConvOp, NumSpatialDimensions> const& problem_shape) {
|
||||
using cute::for_each;
|
||||
using cute::make_seq;
|
||||
|
||||
cute::array<int, NumSpatialDimensions> lower{};
|
||||
if constexpr (ConvOp == conv::Operator::kFprop ||
|
||||
ConvOp == conv::Operator::kWgrad) {
|
||||
for_each(make_seq<NumSpatialDimensions>{}, [&](auto i) {
|
||||
lower[NumSpatialDimensions-1-i] = -1 * problem_shape.lower_padding[i];
|
||||
});
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kDgrad) {
|
||||
for_each(make_seq<NumSpatialDimensions>{}, [&](auto i) {
|
||||
lower[NumSpatialDimensions-1-i] = problem_shape.lower_padding[i] -
|
||||
(problem_shape.shape_B[i+1] - 1) * problem_shape.dilation[i];
|
||||
});
|
||||
}
|
||||
return lower;
|
||||
}
|
||||
|
||||
// Computes the upper/far corner, returning it as a cute::array in [W,H,D] order
|
||||
template <conv::Operator ConvOp, int NumSpatialDimensions>
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
compute_upper_corner_whd(ConvProblemShape<ConvOp, NumSpatialDimensions> const& problem_shape) {
|
||||
using cute::for_each;
|
||||
using cute::make_seq;
|
||||
|
||||
cute::array<int, NumSpatialDimensions> upper{};
|
||||
if constexpr (ConvOp == conv::Operator::kFprop) {
|
||||
for_each(make_seq<NumSpatialDimensions>{}, [&](auto i) {
|
||||
upper[NumSpatialDimensions-1-i] = problem_shape.upper_padding[i] -
|
||||
(problem_shape.shape_B[i+1] - 1) * problem_shape.dilation[i];
|
||||
});
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
for_each(make_seq<NumSpatialDimensions>{}, [&](auto i) {
|
||||
upper[NumSpatialDimensions-1-i] = problem_shape.upper_padding[i] -
|
||||
(problem_shape.shape_C[i+1] - 1) * problem_shape.dilation[i];
|
||||
});
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kDgrad) {
|
||||
for_each(make_seq<NumSpatialDimensions>{}, [&](auto i) {
|
||||
upper[NumSpatialDimensions-1-i] = problem_shape.lower_padding[i] -
|
||||
(problem_shape.shape_B[i+1] - 1) * problem_shape.dilation[i] + problem_shape.shape_C[i+1] - problem_shape.shape_A[i+1];
|
||||
});
|
||||
}
|
||||
return upper;
|
||||
}
|
||||
|
||||
// Compute the lower/near corner of (t,r,s), returning it as a cute::array in [S,R,T] order
|
||||
template <conv::Operator ConvOp, int NumSpatialDimensions>
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
compute_lower_srt(ConvProblemShape<ConvOp, NumSpatialDimensions> const& problem_shape) {
|
||||
using cute::for_each;
|
||||
using cute::make_seq;
|
||||
|
||||
cute::array<int, NumSpatialDimensions> lower{};
|
||||
if constexpr (ConvOp == conv::Operator::kFprop ||
|
||||
ConvOp == conv::Operator::kWgrad) {
|
||||
for_each(make_seq<NumSpatialDimensions>{}, [&](auto i) {
|
||||
lower[NumSpatialDimensions-1-i] = 0;
|
||||
});
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kDgrad) {
|
||||
for_each(make_seq<NumSpatialDimensions>{}, [&](auto i) {
|
||||
lower[NumSpatialDimensions-1-i] = (problem_shape.shape_B[i+1] - 1) * problem_shape.dilation[i];
|
||||
});
|
||||
}
|
||||
return lower;
|
||||
}
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::collective::detail
|
||||
@@ -0,0 +1,616 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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 "cute/tensor_predicate.hpp"
|
||||
#include "cute/arch/cluster_sm90.hpp"
|
||||
#include "cute/arch/copy_sm90.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cute/atom/copy_traits_sm90_im2col.hpp"
|
||||
#include "cute/numeric/arithmetic_tuple.hpp"
|
||||
#include "cute/algorithm/functional.hpp"
|
||||
#include "cute/algorithm/gemm.hpp"
|
||||
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/convnd_problem_shape.hpp"
|
||||
#include "cutlass/conv/dispatch_policy.hpp"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/util/packed_stride.hpp"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv::collective {
|
||||
using namespace cute;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
conv::Operator ConvOp,
|
||||
int Stages,
|
||||
int NumSpatialDims,
|
||||
class ClusterShape,
|
||||
class KernelSchedule,
|
||||
int PipelineAsyncMmaStages,
|
||||
class TileShape_,
|
||||
class ElementA_,
|
||||
class ElementB_,
|
||||
class TiledMma_,
|
||||
class TileTraitsA_,
|
||||
class TileTraitsB_>
|
||||
struct CollectiveConv<
|
||||
MainloopSm90TmaGmmaWarpSpecializedImplicitGemm<
|
||||
ConvOp, Stages, NumSpatialDims, ClusterShape, KernelSchedule, PipelineAsyncMmaStages>,
|
||||
TileShape_,
|
||||
ElementA_,
|
||||
ElementB_,
|
||||
TiledMma_,
|
||||
TileTraitsA_,
|
||||
TileTraitsB_>
|
||||
{
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
using DispatchPolicy = MainloopSm90TmaGmmaWarpSpecializedImplicitGemm<
|
||||
ConvOp, Stages, NumSpatialDims, ClusterShape, KernelSchedule, PipelineAsyncMmaStages>;
|
||||
using TileShape = TileShape_;
|
||||
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 SmemLayoutA = typename TileTraitsA_::SmemLayout;
|
||||
using SmemLayoutB = typename TileTraitsB_::SmemLayout;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
static constexpr int NumSpatialDimensions = DispatchPolicy::NumSpatialDimensions;
|
||||
static constexpr int NumTensorDimensions = NumSpatialDimensions + 2;
|
||||
// Deduce the kernel-facing stride tuple types based on the dispatch policy
|
||||
// (which is a function of the number of spatial dimensions, the algorithm, etc.)
|
||||
using StrideA = decltype(detail::sm90_dispatch_policy_to_stride_A<DispatchPolicy>());
|
||||
using StrideB = decltype(detail::sm90_dispatch_policy_to_stride_B<DispatchPolicy>());
|
||||
|
||||
using MainloopPipeline = cutlass::PipelineTmaAsync<DispatchPolicy::Stages>;
|
||||
|
||||
using PipelineParams = typename MainloopPipeline::Params;
|
||||
using PipelineState = typename cutlass::PipelineState<DispatchPolicy::Stages>;
|
||||
|
||||
// TODO: move pipeline mode tiling into the collective setup phase instead
|
||||
static_assert(rank(SmemLayoutA{}) == 3, "SmemLayout must be rank 3 (M/N, K, PIPE)");
|
||||
static_assert((size<0>(TileShape{}) == size<0>(SmemLayoutA{})), "SmemLayout must be compatible with the tile shape.");
|
||||
static_assert((size<2>(TileShape{}) == size<1>(SmemLayoutA{})), "SmemLayout must be compatible with the tile shape.");
|
||||
|
||||
static_assert(rank(SmemLayoutB{}) == 3, "SmemLayout must be rank 3 (M/N, K, PIPE)");
|
||||
static_assert((size<1>(TileShape{}) == size<0>(SmemLayoutB{})), "SmemLayout must be compatible with the tile shape.");
|
||||
static_assert((size<2>(TileShape{}) == size<1>(SmemLayoutB{})), "SmemLayout must be compatible with the tile shape.");
|
||||
|
||||
static_assert(DispatchPolicy::Stages >= 2, "Specialization requires Stages set to value 1 or more.");
|
||||
static_assert(cute::is_base_of<cute::GMMA::DescriptorIterator, typename TiledMma::FrgTypeA>::value &&
|
||||
cute::is_base_of<cute::GMMA::DescriptorIterator, typename TiledMma::FrgTypeB>::value,
|
||||
"MMA atom must source both A and B operand from smem_desc for this mainloop.");
|
||||
|
||||
// The tma load mode of wgrad is tiled for tensor A and im2col for tensor B while the tma load mode of fprop and dgrad
|
||||
// kernel is im2col for tensor A and tiled for tensor B.
|
||||
static_assert((ConvOp == conv::Operator::kWgrad
|
||||
&& (cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD> || cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD_MULTICAST>))
|
||||
|| (ConvOp != conv::Operator::kWgrad
|
||||
&& (cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD_IM2COL> || cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD_IM2COL_MULTICAST>)),
|
||||
"GmemTiledCopyA - invalid SM90 TMA copy atom specified.");
|
||||
static_assert((ConvOp == conv::Operator::kWgrad
|
||||
&& (cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_IM2COL> || cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_IM2COL_MULTICAST>))
|
||||
|| (ConvOp != conv::Operator::kWgrad
|
||||
&& (cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD> || cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_MULTICAST>)),
|
||||
"GmemTiledCopyB - invalid SM90 TMA copy atom specified.");
|
||||
|
||||
// TMA converts f32 input to tf32 when copying from GMEM to SMEM
|
||||
// For all other types, cast to size equivalent uint type to avoid any rounding by TMA.
|
||||
static constexpr bool ConvertF32toTF32A = cute::is_same_v<float, ElementA>;
|
||||
static constexpr bool ConvertF32toTF32B = cute::is_same_v<float, ElementB>;
|
||||
using InternalElementA = cute::conditional_t<ConvertF32toTF32A, tfloat32_t, uint_bit_t<sizeof_bits_v<ElementA>>>;
|
||||
using InternalElementB = cute::conditional_t<ConvertF32toTF32B, tfloat32_t, uint_bit_t<sizeof_bits_v<ElementB>>>;
|
||||
|
||||
struct SharedStorage
|
||||
{
|
||||
struct TensorStorage : cute::aligned_struct<128> {
|
||||
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;
|
||||
|
||||
static constexpr int K_PIPE_MAX = DispatchPolicy::Stages;
|
||||
static constexpr int K_PIPE_MMAS = DispatchPolicy::PipelineAsyncMmaStages;
|
||||
static constexpr uint32_t TmaTransactionBytes =
|
||||
(size<0>(SmemLayoutA{}) * size<1>(SmemLayoutA{}) * static_cast<uint32_t>(sizeof(InternalElementA)))+
|
||||
(size<0>(SmemLayoutB{}) * size<1>(SmemLayoutB{}) * static_cast<uint32_t>(sizeof(InternalElementB)));
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
using ProblemShape = ConvProblemShape<ConvOp, NumSpatialDimensions>;
|
||||
ProblemShape problem_shape{};
|
||||
ElementA const* ptr_A{nullptr};
|
||||
ElementB const* ptr_B{nullptr};
|
||||
};
|
||||
|
||||
private:
|
||||
// Note that for fprop and 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.
|
||||
|
||||
// Get tma_load_a instantce.
|
||||
template <class TensorA>
|
||||
static constexpr auto
|
||||
get_tma_load_a_instance(TensorA const& tensor_a, typename Arguments::ProblemShape const& problem_shape) {
|
||||
if constexpr (ConvOp == conv::Operator::kFprop || ConvOp == conv::Operator::kDgrad) {
|
||||
// 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);
|
||||
|
||||
// The calculation of gbasis strides for dgrad kernel needs perform negate for dilation values.
|
||||
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_copy(
|
||||
GmemTiledCopyA{},
|
||||
tensor_a,
|
||||
SmemLayoutA{}(_,_,_0{}),
|
||||
product_each(shape(SmemLayoutA{}(_,_,_0{}))),
|
||||
size<1>(ClusterShape{}),
|
||||
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 kernel.
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
return make_tma_copy(
|
||||
GmemTiledCopyA{},
|
||||
tensor_a,
|
||||
SmemLayoutA{}(_,_,_0{}),
|
||||
make_shape(shape<0>(TileShape{}), shape<2>(TileShape{})),
|
||||
size<1>(ClusterShape{}));
|
||||
}
|
||||
}
|
||||
|
||||
// Get tma_load_b instantce.
|
||||
template <class TensorB>
|
||||
static constexpr auto
|
||||
get_tma_load_b_instance(TensorB const& tensor_b, typename Arguments::ProblemShape const& problem_shape) {
|
||||
if constexpr (ConvOp == conv::Operator::kFprop || ConvOp == conv::Operator::kDgrad) {
|
||||
return make_tma_copy(
|
||||
GmemTiledCopyB{},
|
||||
tensor_b,
|
||||
SmemLayoutB{}(_,_,_0{}),
|
||||
make_shape(shape<1>(TileShape{}), shape<2>(TileShape{})),
|
||||
size<0>(ClusterShape{}));
|
||||
}
|
||||
// TMA im2col mode for tensor B in wgrad kernel.
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
// 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_copy(
|
||||
GmemTiledCopyB{},
|
||||
tensor_b,
|
||||
SmemLayoutB{}(_,_,_0{}),
|
||||
product_each(shape(SmemLayoutB{}(_,_,_0{}))),
|
||||
size<0>(ClusterShape{}),
|
||||
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)));
|
||||
}
|
||||
}
|
||||
|
||||
public:
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
using _Submode = decltype(take<0,NumTensorDimensions-1>(typename Arguments::ProblemShape::TensorExtent{}));
|
||||
using ProblemShape = cute::conditional_t<DispatchPolicy::ConvOp == conv::Operator::kWgrad,
|
||||
Shape<int, _Submode, _Submode>,
|
||||
Shape<_Submode, int, _Submode>>;
|
||||
|
||||
// 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{}, int(0)))>;
|
||||
|
||||
using TensorShapeB = cute::conditional_t<ConvOp == conv::Operator::kWgrad,
|
||||
decltype(make_shape(int(0), _Submode{})),
|
||||
decltype(repeat_like(StrideB{}, int32_t(0)))>;
|
||||
|
||||
using TMA_A = decltype(get_tma_load_a_instance(
|
||||
make_tensor(
|
||||
make_gmem_ptr(static_cast<InternalElementA const*>(nullptr)),
|
||||
make_layout(TensorShapeA{}, StrideA{})),
|
||||
ConvProblemShape<ConvOp, NumSpatialDimensions>{}));
|
||||
|
||||
using TMA_B = decltype(get_tma_load_b_instance(
|
||||
make_tensor(
|
||||
make_gmem_ptr(static_cast<InternalElementB const*>(nullptr)),
|
||||
make_layout(TensorShapeB{}, StrideB{})),
|
||||
ConvProblemShape<ConvOp, NumSpatialDimensions>{}));
|
||||
|
||||
// Members
|
||||
TMA_A tma_load_a;
|
||||
TMA_B tma_load_b;
|
||||
ProblemShape problem_shape;
|
||||
};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
// Lowers the host side user facing arguments to the kernel facing lauch params
|
||||
static constexpr Params
|
||||
to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
(void) workspace;
|
||||
// from the flat problem shape arrays of ConvProblemShape<ConvOp, N>, create a rank-3 MNK problem shape tuple
|
||||
// tma desc creation depends on the original untransformed domain.
|
||||
|
||||
// A extents.
|
||||
auto shape_A_orig = args.problem_shape.get_shape_A();
|
||||
// B extents.
|
||||
auto shape_B_orig = args.problem_shape.get_shape_B();
|
||||
|
||||
// Fill inferred cute strides from flat stride arrays
|
||||
auto dA = make_cute_packed_stride(StrideA{}, args.problem_shape.stride_A, ConvOp);
|
||||
auto dB = make_cute_packed_stride(StrideB{}, args.problem_shape.stride_B, ConvOp);
|
||||
|
||||
auto ptr_A = reinterpret_cast<InternalElementA const*>(args.ptr_A);
|
||||
auto ptr_B = reinterpret_cast<InternalElementB const*>(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 tma_load_a = get_tma_load_a_instance(tensor_a, args.problem_shape);
|
||||
auto tma_load_b = get_tma_load_b_instance(tensor_b, args.problem_shape);
|
||||
|
||||
auto problem_shape_mnk = args.problem_shape.get_transformed_problem_shape_MNK();
|
||||
|
||||
return {
|
||||
tma_load_a,
|
||||
tma_load_b,
|
||||
problem_shape_mnk
|
||||
};
|
||||
}
|
||||
|
||||
template<class ProblemShape>
|
||||
CUTLASS_HOST_DEVICE 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
|
||||
implementable &= args.problem_shape.stride_A[NumTensorDimensions-1] == 1;
|
||||
implementable &= args.problem_shape.stride_B[NumTensorDimensions-1] == 1;
|
||||
|
||||
constexpr int tma_alignment_bits = 128;
|
||||
// A extents.
|
||||
auto shape_A_orig = args.problem_shape.get_shape_A();
|
||||
// B extents.
|
||||
auto shape_B_orig = args.problem_shape.get_shape_B();
|
||||
constexpr int min_tma_aligned_elements_A = tma_alignment_bits / cutlass::sizeof_bits<ElementA>::value;
|
||||
implementable = implementable && cutlass::detail::check_alignment<min_tma_aligned_elements_A>(shape_A_orig, StrideA{});
|
||||
constexpr int min_tma_aligned_elements_B = tma_alignment_bits / cutlass::sizeof_bits<ElementB>::value;
|
||||
implementable = implementable && cutlass::detail::check_alignment<min_tma_aligned_elements_B>(shape_B_orig, StrideB{});
|
||||
|
||||
if (!implementable) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Problem Size doesn't meet the minimum alignment requirements for TMA.\n");
|
||||
return false;
|
||||
}
|
||||
|
||||
// Check valid padding values for TMA_LOAD_IM2COL
|
||||
constexpr int padding_limit = (ProblemShape::RankS == 1) ? 65536 : (ProblemShape::RankS == 2 ? 256 : 16);
|
||||
for (int i = 0; i < problem_shape.RankS; ++i) {
|
||||
implementable = implementable && problem_shape.lower_padding[i] <= padding_limit && problem_shape.lower_padding[i] >= 0;
|
||||
implementable = implementable && problem_shape.upper_padding[i] <= padding_limit && problem_shape.upper_padding[i] >= 0;
|
||||
}
|
||||
|
||||
if (!implementable) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Padding values don't meet requirements for TMA LOAD IM2COL.\n");
|
||||
return false;
|
||||
}
|
||||
|
||||
if (problem_shape.groups > 1) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: This kernel does not support conv groups > 1.\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) {
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_a.get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(mainloop_params.tma_load_b.get_tma_descriptor());
|
||||
}
|
||||
|
||||
/// Perform a collective-scoped matrix multiply-accumulate
|
||||
/// Producer Perspective
|
||||
template <
|
||||
class TensorA, class TMA_LOAD_A,
|
||||
class TensorB, class TMA_LOAD_B,
|
||||
class KTileIterator
|
||||
>
|
||||
CUTLASS_DEVICE void
|
||||
load(MainloopPipeline pipeline,
|
||||
PipelineState smem_pipe_producer_state,
|
||||
TensorA const& gA, TMA_LOAD_A& tma_load_a,
|
||||
TensorB const& gB, TMA_LOAD_B& tma_load_b,
|
||||
KTileIterator k_tile_iter, int k_tile_count,
|
||||
int thad_idx,
|
||||
TensorStorage& shared_tensors) {
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % 4;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
|
||||
if (warp_idx_in_warp_group == 0 and lane_predicate) {
|
||||
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)
|
||||
|
||||
//
|
||||
// Prepare the TMA loads for A and B
|
||||
//
|
||||
|
||||
dim3 cluster_local_block_id = cute::block_id_in_cluster();
|
||||
auto block_tma_a = tma_load_a.get_slice(cluster_local_block_id.y);
|
||||
auto block_tma_b = tma_load_b.get_slice(cluster_local_block_id.x);
|
||||
|
||||
// Applies the mapping from block_tma_a
|
||||
Tensor tAgA = block_tma_a.partition_S(gA); // (TMA,TMA_M,TMA_K,k)
|
||||
Tensor tAsA = block_tma_a.partition_D(sA); // (TMA,TMA_M,TMA_K,PIPE)
|
||||
|
||||
Tensor tBgB = block_tma_b.partition_S(gB); // (TMA,TMA_N,TMA_K,k)
|
||||
Tensor tBsB = block_tma_b.partition_D(sB); // (TMA,TMA_N,TMA_K,PIPE)
|
||||
|
||||
uint16_t mcast_mask_a = 0;
|
||||
uint16_t mcast_mask_b = 0;
|
||||
|
||||
// Issue TmaLoads
|
||||
// Maps the tile -> block, value
|
||||
if constexpr (cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD_IM2COL_MULTICAST> ||
|
||||
cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD_MULTICAST>) {
|
||||
auto block_layout = Layout<typename DispatchPolicy::ClusterShape>{}; // (m,n) -> block_id
|
||||
for (int n = 0; n < size<1>(block_layout); ++n) {
|
||||
mcast_mask_a |= (uint16_t(1) << block_layout(cluster_local_block_id.x,n,Int<0>{}));
|
||||
}
|
||||
}
|
||||
|
||||
if constexpr (cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_IM2COL_MULTICAST> ||
|
||||
cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_MULTICAST>) {
|
||||
auto block_layout = Layout<typename DispatchPolicy::ClusterShape>{}; // (m,n) -> block_id
|
||||
for (int m = 0; m < size<0>(block_layout); ++m) {
|
||||
mcast_mask_b |= (uint16_t(1) << block_layout(m,cluster_local_block_id.y,Int<0>{}));
|
||||
}
|
||||
}
|
||||
|
||||
// Mainloop
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for ( ; k_tile_count > 0; --k_tile_count) {
|
||||
// LOCK smem_pipe_producer_state for _writing_
|
||||
pipeline.producer_acquire(smem_pipe_producer_state);
|
||||
|
||||
//
|
||||
// Copy gmem to smem for *k_tile_iter
|
||||
//
|
||||
|
||||
using BarrierType = typename MainloopPipeline::ProducerBarrierType;
|
||||
BarrierType* tma_barrier = pipeline.producer_get_barrier(smem_pipe_producer_state);
|
||||
|
||||
int write_stage = smem_pipe_producer_state.index();
|
||||
copy(tma_load_a.with(*tma_barrier, mcast_mask_a), tAgA(_,_,_,*k_tile_iter), tAsA(_,_,_,write_stage));
|
||||
copy(tma_load_b.with(*tma_barrier, mcast_mask_b), tBgB(_,_,_,*k_tile_iter), tBsB(_,_,_,write_stage));
|
||||
++k_tile_iter;
|
||||
|
||||
// Advance smem_pipe_producer_state
|
||||
++smem_pipe_producer_state;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Perform a Producer Epilogue to prevent early exit of blocks in a Cluster
|
||||
CUTLASS_DEVICE void
|
||||
load_tail(MainloopPipeline pipeline, PipelineState smem_pipe_producer_state) {
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_idx_in_warp_group = warp_idx % 4;
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
|
||||
// Issue the epilogue waits
|
||||
if (warp_idx_in_warp_group == 0 and lane_predicate) {
|
||||
/* This helps avoid early exit of blocks 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(smem_pipe_producer_state);
|
||||
}
|
||||
}
|
||||
|
||||
/// Perform a collective-scoped matrix multiply-accumulate
|
||||
/// Consumer Perspective
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE void
|
||||
mma(MainloopPipeline pipeline,
|
||||
PipelineState smem_pipe_consumer_state,
|
||||
FrgTensorC& accum,
|
||||
int k_tile_count,
|
||||
int thread_idx,
|
||||
TensorStorage& shared_tensors,
|
||||
Params const& mainloop_params) {
|
||||
static_assert(is_rmem<FrgTensorC>::value, "C tensor must be rmem resident.");
|
||||
|
||||
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)
|
||||
|
||||
//
|
||||
// Define C accumulators and A/B partitioning
|
||||
//
|
||||
|
||||
TiledMma tiled_mma;
|
||||
auto thread_mma = tiled_mma.get_thread_slice(thread_idx);
|
||||
|
||||
Tensor tCsA = thread_mma.partition_A(sA); // (MMA,MMA_M,MMA_K,PIPE)
|
||||
Tensor tCsB = thread_mma.partition_B(sB); // (MMA,MMA_N,MMA_K,PIPE)
|
||||
|
||||
// Allocate "fragments/descriptors"
|
||||
Tensor tCrA = thread_mma.make_fragment_A(tCsA); // (MMA,MMA_M,MMA_K,PIPE)
|
||||
Tensor tCrB = thread_mma.make_fragment_B(tCsB); // (MMA,MMA_N,MMA_K,PIPE)
|
||||
|
||||
CUTE_STATIC_ASSERT_V(size<1>(tCsA) == size<1>(accum)); // M
|
||||
CUTE_STATIC_ASSERT_V(size<1>(tCsB) == size<2>(accum)); // N
|
||||
CUTE_STATIC_ASSERT_V(size<2>(tCsA) == size<2>(tCsB)); // K
|
||||
CUTE_STATIC_ASSERT_V(size<3>(tCsA) == size<3>(tCsB)); // PIPE
|
||||
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<2>(sA)); // PIPE
|
||||
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<2>(sB)); // PIPE
|
||||
|
||||
//
|
||||
// PIPELINED MAIN LOOP
|
||||
//
|
||||
static_assert((0 <= K_PIPE_MMAS) && (K_PIPE_MMAS < K_PIPE_MAX),
|
||||
"ERROR : Incorrect number of MMAs in flight");
|
||||
|
||||
// We release buffers to producer warps(dma load) with some mmas in flight
|
||||
PipelineState smem_pipe_release = smem_pipe_consumer_state;
|
||||
|
||||
// Prologue GMMAs
|
||||
int prologue_mma_count = min(K_PIPE_MMAS, k_tile_count);
|
||||
|
||||
tiled_mma.accumulate_ = GMMA::ScaleOut::Zero;
|
||||
|
||||
warpgroup_fence_operand(accum);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_tile_prologue = prologue_mma_count; k_tile_prologue > 0; --k_tile_prologue) {
|
||||
// WAIT on smem_pipe_consumer_state until its data are available (phase bit flips from rdPhaseBit value)
|
||||
pipeline.consumer_wait(smem_pipe_consumer_state);
|
||||
|
||||
int read_stage = smem_pipe_consumer_state.index();
|
||||
warpgroup_arrive();
|
||||
// Unroll the K mode manually to set scale D 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), accum);
|
||||
tiled_mma.accumulate_ = GMMA::ScaleOut::One;
|
||||
}
|
||||
|
||||
warpgroup_commit_batch();
|
||||
|
||||
++smem_pipe_consumer_state;
|
||||
}
|
||||
|
||||
warpgroup_fence_operand(accum);
|
||||
// Mainloop GMMAs
|
||||
k_tile_count -= prologue_mma_count;
|
||||
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for ( ; k_tile_count > 0; --k_tile_count) {
|
||||
// WAIT on smem_pipe_consumer_state until its data are available (phase bit flips from rdPhaseBit value)
|
||||
pipeline.consumer_wait(smem_pipe_consumer_state);
|
||||
|
||||
//
|
||||
// Compute on k_tile
|
||||
//
|
||||
|
||||
int read_stage = smem_pipe_consumer_state.index();
|
||||
warpgroup_fence_operand(accum);
|
||||
warpgroup_arrive();
|
||||
// Unroll the K mode manually to set scale D to 1
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
|
||||
// (V,M) x (V,N) => (V,M,N)
|
||||
cute::gemm(tiled_mma, tCrA(_,_,k_block,read_stage), tCrB(_,_,k_block,read_stage), accum);
|
||||
tiled_mma.accumulate_ = GMMA::ScaleOut::One;
|
||||
}
|
||||
warpgroup_commit_batch();
|
||||
|
||||
/// Wait on the GMMA barrier for K_PIPE_MMAS (or fewer) outstanding to ensure smem_pipe_producer_state is consumed
|
||||
warpgroup_wait<K_PIPE_MMAS>();
|
||||
warpgroup_fence_operand(accum);
|
||||
|
||||
// UNLOCK smem_pipe_release, done _computing_ on it
|
||||
pipeline.consumer_release(smem_pipe_release);
|
||||
|
||||
// Advance smem_pipe_consumer_state and smem_pipe_release
|
||||
++smem_pipe_consumer_state;
|
||||
++smem_pipe_release;
|
||||
}
|
||||
|
||||
warpgroup_fence_operand(accum);
|
||||
}
|
||||
|
||||
/// Perform a Consumer Epilogue to release all buffers
|
||||
CUTLASS_DEVICE void
|
||||
mma_tail(MainloopPipeline pipeline, PipelineState smem_pipe_release, int k_tile_count) {
|
||||
// Prologue GMMAs
|
||||
int prologue_mma_count = min(K_PIPE_MMAS, k_tile_count);
|
||||
k_tile_count -= prologue_mma_count;
|
||||
|
||||
smem_pipe_release.advance(k_tile_count);
|
||||
|
||||
// Wait on all GMMAs to complete
|
||||
warpgroup_wait<0>();
|
||||
|
||||
for (int count = 0; count < prologue_mma_count; ++count) {
|
||||
pipeline.consumer_release(smem_pipe_release); // UNLOCK smem_pipe_release, done _computing_ on it
|
||||
++smem_pipe_release;
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -35,7 +35,7 @@
|
||||
activation (NHWC),
|
||||
filter (KRSC),
|
||||
output (NPQK),
|
||||
pading (pad_h, pad_w),
|
||||
pading (pad_h, pad_w),
|
||||
stride (stride_h, stride_w),
|
||||
dilation (dilation_h, dilation_w).
|
||||
|
||||
@@ -109,7 +109,7 @@ public:
|
||||
Mode mode
|
||||
):
|
||||
N(N), H(H), W(W), C(C), P(P), Q(Q), K(K), R(R), S(S),
|
||||
pad_h(R / 2), pad_w(S / 2), stride_h(1), stride_w(1), dilation_h(1), dilation_w(1),
|
||||
pad_h(R / 2), pad_w(S / 2), stride_h(1), stride_w(1), dilation_h(1), dilation_w(1),
|
||||
mode(mode), split_k_slices(1), groups (1) { }
|
||||
|
||||
/// Constructor
|
||||
@@ -133,9 +133,9 @@ public:
|
||||
Mode mode,
|
||||
int split_k_slices = 1,
|
||||
int groups = 1
|
||||
):
|
||||
):
|
||||
N(N), H(H), W(W), C(C), P(P), Q(Q), K(K), R(R), S(S),
|
||||
pad_h(pad_h), pad_w(pad_w), stride_h(stride_h), stride_w(stride_w),
|
||||
pad_h(pad_h), pad_w(pad_w), stride_h(stride_h), stride_w(stride_w),
|
||||
dilation_h(dilation_h), dilation_w(dilation_w),
|
||||
mode(mode), split_k_slices(split_k_slices), groups (groups) { }
|
||||
|
||||
@@ -156,8 +156,8 @@ public:
|
||||
N(input_size.n()), H(input_size.h()), W(input_size.w()), C(input_size.c()),
|
||||
P(output_size.h()), Q(output_size.w()),
|
||||
K(filter_size.n()), R(filter_size.h()), S(filter_size.w()),
|
||||
pad_h(padding[0]), pad_w(padding[2]),
|
||||
stride_h(stride.row()), stride_w(stride.column()),
|
||||
pad_h(padding[0]), pad_w(padding[2]),
|
||||
stride_h(stride.row()), stride_w(stride.column()),
|
||||
dilation_h(dilation.row()), dilation_w(dilation.column()),
|
||||
mode(mode), split_k_slices(split_k_slices), groups(groups) {}
|
||||
|
||||
@@ -167,7 +167,7 @@ public:
|
||||
Conv2dProblemSize(
|
||||
cutlass::Tensor4DCoord input_size, // NHWC
|
||||
cutlass::Tensor4DCoord filter_size, // KRSC
|
||||
cutlass::Tensor4DCoord padding, // pad_h, _, pad_w, _
|
||||
cutlass::Tensor4DCoord padding, // pad_h, upper_pad_h, pad_w, upper_pad_w
|
||||
cutlass::MatrixCoord stride, // stride_h, stride_w
|
||||
cutlass::MatrixCoord dilation, // dilation_h, dilation_w
|
||||
cutlass::conv::Mode mode = cutlass::conv::Mode::kCrossCorrelation,
|
||||
@@ -177,12 +177,12 @@ public:
|
||||
N(input_size.n()), H(input_size.h()), W(input_size.w()), C(input_size.c()),
|
||||
K(filter_size.n()), R(filter_size.h()), S(filter_size.w()),
|
||||
pad_h(padding[0]), pad_w(padding[2]),
|
||||
stride_h(stride.row()), stride_w(stride.column()),
|
||||
stride_h(stride.row()), stride_w(stride.column()),
|
||||
dilation_h(dilation.row()), dilation_w(dilation.column()),
|
||||
mode(mode), split_k_slices(split_k_slices), groups(groups) {
|
||||
// set output P and Q
|
||||
P = ((H + pad_h * 2 - R * dilation_h) / stride_h) + 1;
|
||||
Q = ((W + pad_w * 2 - S * dilation_w) / stride_w) + 1;
|
||||
P = ((H + pad_h + padding[1] - R * dilation_h) / stride_h) + 1;
|
||||
Q = ((W + pad_w + padding[3] - S * dilation_w) / stride_w) + 1;
|
||||
}
|
||||
|
||||
/// Constructs convolution problem size from cutlass Tensor4DCoord and MatrixCoord
|
||||
@@ -199,7 +199,7 @@ public:
|
||||
N(input_size.n()), H(input_size.h()), W(input_size.w()), C(input_size.c()),
|
||||
P(output_size.h()), Q(output_size.w()),
|
||||
K(filter_size.n()), R(filter_size.h()), S(filter_size.w()),
|
||||
pad_h(R / 2), pad_w(S / 2), stride_h(1), stride_w(1),
|
||||
pad_h(R / 2), pad_w(S / 2), stride_h(1), stride_w(1),
|
||||
dilation_h(1), dilation_w(1),
|
||||
mode(mode), split_k_slices(split_k_slices), groups(groups) {}
|
||||
|
||||
|
||||
@@ -91,7 +91,7 @@ public:
|
||||
Conv3dProblemSize():
|
||||
Conv2dProblemSize(),
|
||||
D(0), T(0), Z(0),
|
||||
pad_d(0),
|
||||
pad_d(0),
|
||||
stride_d(1),
|
||||
dilation_d(1) { }
|
||||
|
||||
@@ -205,6 +205,34 @@ public:
|
||||
Z = ((D + pad_d * 2 - T * dilation_d) / stride_d) + 1;
|
||||
}
|
||||
|
||||
/// Constructs convolution problem size from cutlass Tensor5DCoord, Coord3D
|
||||
// *computes* output size and sets Z, P and Q (include all data members in ctor)
|
||||
CUTLASS_HOST_DEVICE
|
||||
Conv3dProblemSize(
|
||||
cutlass::Tensor5DCoord input_size, // NDHWC
|
||||
cutlass::Tensor5DCoord filter_size, // KTRSC
|
||||
CUTLASS_STL_NAMESPACE::tuple<Coord3D, Coord3D> padding, // Coord3D {pad_d, pad_h, pad_w} & Coord3D {far pad_d, pad_h, pad_w} to calculate o/p/q
|
||||
Coord3D stride, // stride_d, stride_h, stride_w
|
||||
Coord3D dilation, // dilation_d, dilation_h, dilation_w
|
||||
cutlass::conv::Mode mode = cutlass::conv::Mode::kCrossCorrelation,
|
||||
int split_k_slices = 1,
|
||||
int groups = 1
|
||||
):
|
||||
Conv2dProblemSize(
|
||||
{input_size.n(), input_size.h(), input_size.w(), input_size.c()},
|
||||
{filter_size.n(), filter_size.h(), filter_size.w(), filter_size.c()},
|
||||
{CUTLASS_STL_NAMESPACE::get<0>(padding)[1], CUTLASS_STL_NAMESPACE::get<1>(padding)[1],
|
||||
CUTLASS_STL_NAMESPACE::get<0>(padding)[2], CUTLASS_STL_NAMESPACE::get<1>(padding)[2]},
|
||||
{stride[1], stride[2]},
|
||||
{dilation[1], dilation[2]},
|
||||
mode, split_k_slices, groups),
|
||||
D(input_size.d()), T(filter_size.d()),
|
||||
pad_d(CUTLASS_STL_NAMESPACE::get<0>(padding)[0]), stride_d(stride[0]), dilation_d(dilation[0])
|
||||
{
|
||||
// set output Z
|
||||
Z = ((D + pad_d + CUTLASS_STL_NAMESPACE::get<1>(padding)[0] - T * dilation_d) / stride_d) + 1;
|
||||
}
|
||||
|
||||
/// Equality operator (ignores mode and split_k_slice)
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool operator==(Conv3dProblemSize const &conv) const {
|
||||
@@ -282,7 +310,7 @@ public:
|
||||
return (N * Z * P * Q * K);
|
||||
}
|
||||
|
||||
/// Returns output extent as Tensor5DCoord
|
||||
/// Returns padding as Coord3D
|
||||
CUTLASS_HOST_DEVICE
|
||||
Coord3D padding() const {
|
||||
|
||||
|
||||
@@ -0,0 +1,574 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
/*! \file
|
||||
\brief This file contains definitions and utility functions for describing convolution problem shapes.
|
||||
*/
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/tensor_coord.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
|
||||
#include "cute/container/array.hpp"
|
||||
|
||||
#if ! defined(__CUDACC_RTC__)
|
||||
#include <initializer_list>
|
||||
#endif
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Implements the user facing argument for all CUTLASS 3.x convolutions in a rank agnostic fashion.
|
||||
// All tensors are flat and by default treated as layout right (NDHWC, KTRSC, NZPQK)
|
||||
// Supports asymmetric padding, traversal strides, dilations, and all conv algorithm types.
|
||||
template <
|
||||
conv::Operator ConvOp_,
|
||||
int NumSpatialDimensions
|
||||
>
|
||||
struct ConvProblemShape {
|
||||
//
|
||||
// Alias types for members
|
||||
//
|
||||
static constexpr int RankS = NumSpatialDimensions;
|
||||
static constexpr int RankT = NumSpatialDimensions + 2;
|
||||
static constexpr conv::Operator ConvOp = ConvOp_;
|
||||
using SpatialExtent = cute::array<int, RankS>;
|
||||
using TensorExtent = cute::array<int, RankT>;
|
||||
using TensorStride = cute::array<int64_t, RankT>;
|
||||
using ShapePadding = SpatialExtent;
|
||||
using TraversalStride = SpatialExtent;
|
||||
using ShapeDilation = SpatialExtent;
|
||||
using Corner = SpatialExtent;
|
||||
|
||||
//
|
||||
// Members
|
||||
//
|
||||
cutlass::conv::Mode mode{};
|
||||
TensorExtent shape_A{};
|
||||
TensorStride stride_A{};
|
||||
TensorExtent shape_B{};
|
||||
TensorStride stride_B{};
|
||||
TensorExtent shape_C{};
|
||||
TensorStride stride_C{};
|
||||
|
||||
// asymmetric padding, both upper and lower padding must be >= 0
|
||||
ShapePadding lower_padding{};
|
||||
ShapePadding upper_padding{};
|
||||
TraversalStride traversal_stride{};
|
||||
ShapeDilation dilation{};
|
||||
int groups = 1;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
ConvProblemShape() = default;
|
||||
|
||||
// Constructor accepts user facing arguments and computes to stores the corners as its internal state
|
||||
ConvProblemShape(
|
||||
conv::Mode mode, // convolution/cross-correlation
|
||||
TensorExtent shape_act, // [n,d,h,w,c]
|
||||
TensorStride stride_act, // [n,d,h,w,c]
|
||||
TensorExtent shape_flt, // [k,t,r,s,c]
|
||||
TensorStride stride_flt, // [k,t,r,s,c]
|
||||
ShapePadding lower_padding, // [pad_d, pad_h, pad_w]
|
||||
ShapePadding upper_padding, // [pad_d, pad_h, pad_w]
|
||||
TraversalStride tstride, // [stride_d, stride_h, stride_w]
|
||||
ShapeDilation dilation, // [dilation_d, dilation_h, dilation_w]
|
||||
int groups)
|
||||
: mode(mode)
|
||||
, lower_padding(lower_padding)
|
||||
, upper_padding(upper_padding)
|
||||
, traversal_stride(tstride)
|
||||
, dilation(dilation)
|
||||
, groups(groups) {
|
||||
|
||||
auto [shape_xformed_act, stride_xformed_act] = calculate_xformed_act(shape_act, shape_flt);
|
||||
set_shape_stride_ABC(shape_act, stride_act, shape_flt, stride_flt, shape_xformed_act, stride_xformed_act);
|
||||
}
|
||||
|
||||
// Allow user input of xformed activation stride to support non-packed strides.
|
||||
ConvProblemShape(
|
||||
conv::Mode mode, // convolution/cross-correlation
|
||||
TensorExtent shape_act, // [n,d,h,w,c]
|
||||
TensorStride stride_act, // [n,d,h,w,c]
|
||||
TensorExtent shape_flt, // [k,t,r,s,c]
|
||||
TensorStride stride_flt, // [k,t,r,s,c]
|
||||
TensorStride stride_xformed_act, // [n,z,p,q,k]
|
||||
ShapePadding lower_padding, // [pad_d, pad_h, pad_w]
|
||||
ShapePadding upper_padding, // [pad_d, pad_h, pad_w]
|
||||
TraversalStride tstride, // [stride_d, stride_h, stride_w]
|
||||
ShapeDilation dilation, // [dilation_d, dilation_h, dilation_w]
|
||||
int groups)
|
||||
: mode(mode)
|
||||
, lower_padding(lower_padding)
|
||||
, upper_padding(upper_padding)
|
||||
, traversal_stride(tstride)
|
||||
, dilation(dilation)
|
||||
, groups(groups) {
|
||||
|
||||
CUTLASS_ASSERT(stride_act[RankT - 1] == 1);
|
||||
CUTLASS_ASSERT(stride_flt[RankT - 1] == 1);
|
||||
CUTLASS_ASSERT(stride_xformed_act[RankT - 1] == 1);
|
||||
|
||||
auto stride_act_packed = packed_stride_right_major(shape_act);
|
||||
auto stride_flt_packed = packed_stride_right_major(shape_flt);
|
||||
auto [shape_xformed_act, stride_xformed_act_packed] = calculate_xformed_act(shape_act, shape_flt);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i = 0; i < RankT - 1; ++i) {
|
||||
CUTLASS_ASSERT(stride_act[i] >= stride_act_packed[i]);
|
||||
CUTLASS_ASSERT(stride_flt[i] >= stride_flt_packed[i]);
|
||||
CUTLASS_ASSERT(stride_xformed_act[i] >= stride_xformed_act_packed[i]);
|
||||
}
|
||||
|
||||
set_shape_stride_ABC(shape_act, stride_act, shape_flt, stride_flt, shape_xformed_act, stride_xformed_act);
|
||||
}
|
||||
|
||||
// Constructor accepts user facing arguments and presume packed tensor strides in canonical (CWHDN) order.
|
||||
ConvProblemShape(
|
||||
conv::Mode mode,
|
||||
TensorExtent shape_act,
|
||||
TensorExtent shape_flt,
|
||||
ShapePadding lower_padding,
|
||||
ShapePadding upper_padding,
|
||||
TraversalStride tstride,
|
||||
ShapeDilation dilation,
|
||||
int groups)
|
||||
: ConvProblemShape(
|
||||
mode,
|
||||
shape_act,
|
||||
packed_stride_right_major(shape_act),
|
||||
shape_flt,
|
||||
packed_stride_right_major(shape_flt),
|
||||
lower_padding,
|
||||
upper_padding,
|
||||
tstride,
|
||||
dilation,
|
||||
groups) {
|
||||
}
|
||||
|
||||
#if ! defined(__CUDACC_RTC__)
|
||||
// Constructor accepts user facing arguments and computes to stores the corners as its internal state
|
||||
ConvProblemShape(
|
||||
conv::Mode mode,
|
||||
std::initializer_list<int> shape_act_,
|
||||
std::initializer_list<int64_t> stride_act_,
|
||||
std::initializer_list<int> shape_flt_,
|
||||
std::initializer_list<int64_t> stride_flt_,
|
||||
std::initializer_list<int> lower_padding_,
|
||||
std::initializer_list<int> upper_padding_,
|
||||
std::initializer_list<int> traversal_stride_,
|
||||
std::initializer_list<int> dilation_,
|
||||
int groups)
|
||||
: mode(mode)
|
||||
, groups(groups) {
|
||||
|
||||
TensorExtent shape_act{};
|
||||
TensorStride stride_act{};
|
||||
TensorExtent shape_flt{};
|
||||
TensorStride stride_flt{};
|
||||
|
||||
assert(shape_act_.size() == shape_act.size());
|
||||
assert(stride_act_.size() == stride_act.size());
|
||||
assert(shape_flt_.size() == shape_flt.size());
|
||||
assert(stride_flt_.size() == stride_flt.size());
|
||||
assert(lower_padding_.size() == lower_padding.size());
|
||||
assert(upper_padding_.size() == upper_padding.size());
|
||||
assert(traversal_stride_.size() == traversal_stride.size());
|
||||
assert(dilation_.size() == dilation.size());
|
||||
|
||||
std::copy(shape_act_.begin(), shape_act_.end(), shape_act.begin());
|
||||
std::copy(stride_act_.begin(), stride_act_.end(), stride_act.begin());
|
||||
std::copy(shape_flt_.begin(), shape_flt_.end(), shape_flt.begin());
|
||||
std::copy(stride_flt_.begin(), stride_flt_.end(), stride_flt.begin());
|
||||
std::copy(lower_padding_.begin(), lower_padding_.end(), lower_padding.begin());
|
||||
std::copy(upper_padding_.begin(), upper_padding_.end(), upper_padding.begin());
|
||||
std::copy(traversal_stride_.begin(), traversal_stride_.end(), traversal_stride.begin());
|
||||
std::copy(dilation_.begin(), dilation_.end(), dilation.begin());
|
||||
|
||||
auto [shape_xformed_act, stride_xformed_act] = calculate_xformed_act(shape_act, shape_flt);
|
||||
set_shape_stride_ABC(shape_act, stride_act, shape_flt, stride_flt, shape_xformed_act, stride_xformed_act);
|
||||
}
|
||||
|
||||
// Allow user input of xformed activation stride to support non-packed strides.
|
||||
ConvProblemShape(
|
||||
conv::Mode mode,
|
||||
std::initializer_list<int> shape_act_,
|
||||
std::initializer_list<int64_t> stride_act_,
|
||||
std::initializer_list<int> shape_flt_,
|
||||
std::initializer_list<int64_t> stride_flt_,
|
||||
std::initializer_list<int64_t> stride_xformed_act_,
|
||||
std::initializer_list<int> lower_padding_,
|
||||
std::initializer_list<int> upper_padding_,
|
||||
std::initializer_list<int> traversal_stride_,
|
||||
std::initializer_list<int> dilation_,
|
||||
int groups)
|
||||
: mode(mode)
|
||||
, groups(groups) {
|
||||
TensorExtent shape_act{};
|
||||
TensorStride stride_act{};
|
||||
TensorExtent shape_flt{};
|
||||
TensorStride stride_flt{};
|
||||
TensorStride stride_xformed_act{};
|
||||
|
||||
std::copy(shape_act_.begin(), shape_act_.end(), shape_act.begin());
|
||||
std::copy(stride_act_.begin(), stride_act_.end(), stride_act.begin());
|
||||
std::copy(shape_flt_.begin(), shape_flt_.end(), shape_flt.begin());
|
||||
std::copy(stride_flt_.begin(), stride_flt_.end(), stride_flt.begin());
|
||||
std::copy(stride_xformed_act_.begin(), stride_xformed_act_.end(), stride_xformed_act.begin());
|
||||
std::copy(lower_padding_.begin(), lower_padding_.end(), lower_padding.begin());
|
||||
std::copy(upper_padding_.begin(), upper_padding_.end(), upper_padding.begin());
|
||||
std::copy(traversal_stride_.begin(), traversal_stride_.end(), traversal_stride.begin());
|
||||
std::copy(dilation_.begin(), dilation_.end(), dilation.begin());
|
||||
|
||||
CUTLASS_ASSERT(stride_act[RankT - 1] == 1);
|
||||
CUTLASS_ASSERT(stride_flt[RankT - 1] == 1);
|
||||
CUTLASS_ASSERT(stride_xformed_act[RankT - 1] == 1);
|
||||
|
||||
auto stride_act_packed = packed_stride_right_major(shape_act);
|
||||
auto stride_flt_packed = packed_stride_right_major(shape_flt);
|
||||
auto [shape_xformed_act, stride_xformed_act_packed] = calculate_xformed_act(shape_act, shape_flt);
|
||||
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for(int i = 0; i < RankT - 1; ++i) {
|
||||
CUTLASS_ASSERT(stride_act[i] >= stride_act_packed[i]);
|
||||
CUTLASS_ASSERT(stride_flt[i] >= stride_flt_packed[i]);
|
||||
CUTLASS_ASSERT(stride_xformed_act[i] >= stride_xformed_act_packed[i]);
|
||||
}
|
||||
|
||||
set_shape_stride_ABC(shape_act, stride_act, shape_flt, stride_flt, shape_xformed_act, stride_xformed_act);
|
||||
}
|
||||
|
||||
// Constructor accepts user facing arguments and computes to stores the corners as its internal state
|
||||
ConvProblemShape(
|
||||
conv::Mode mode,
|
||||
std::initializer_list<int> shape_act_,
|
||||
std::initializer_list<int> shape_flt_,
|
||||
std::initializer_list<int> lower_padding_,
|
||||
std::initializer_list<int> upper_padding_,
|
||||
std::initializer_list<int> traversal_stride_,
|
||||
std::initializer_list<int> dilation_,
|
||||
int groups)
|
||||
: mode(mode)
|
||||
, groups(groups) {
|
||||
TensorExtent shape_act{};
|
||||
TensorStride stride_act{};
|
||||
TensorExtent shape_flt{};
|
||||
TensorStride stride_flt{};
|
||||
|
||||
assert(shape_act_.size() == shape_act.size());
|
||||
assert(shape_flt_.size() == shape_flt.size());
|
||||
assert(lower_padding_.size() == lower_padding.size());
|
||||
assert(upper_padding_.size() == upper_padding.size());
|
||||
assert(traversal_stride_.size() == traversal_stride.size());
|
||||
assert(dilation_.size() == dilation.size());
|
||||
|
||||
std::copy(shape_act_.begin(), shape_act_.end(), shape_act.begin());
|
||||
std::copy(shape_flt_.begin(), shape_flt_.end(), shape_flt.begin());
|
||||
std::copy(lower_padding_.begin(), lower_padding_.end(), lower_padding.begin());
|
||||
std::copy(upper_padding_.begin(), upper_padding_.end(), upper_padding.begin());
|
||||
std::copy(traversal_stride_.begin(), traversal_stride_.end(), traversal_stride.begin());
|
||||
std::copy(dilation_.begin(), dilation_.end(), dilation.begin());
|
||||
stride_act = packed_stride_right_major(shape_act);
|
||||
stride_flt = packed_stride_right_major(shape_flt);
|
||||
|
||||
auto [shape_xformed_act, stride_xformed_act] = calculate_xformed_act(shape_act, shape_flt);
|
||||
set_shape_stride_ABC(shape_act, stride_act, shape_flt, stride_flt, shape_xformed_act, stride_xformed_act);
|
||||
}
|
||||
#endif // not defined(__CUDACC_RTC__)
|
||||
|
||||
// Set shape and stride of tensor A/B/C according to following table:
|
||||
// | | Fprop | Dgrad | Wgrad |
|
||||
// | ------ | ------ | ------ | ------|
|
||||
// | ShapeA | NDHWC | NZPQK | NZPQK |
|
||||
// | ShapeB | KTRSC | KTRSC | NDHWC |
|
||||
// | ShapeC | NZPQK | NDHWC | KTRSC |
|
||||
//
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr void
|
||||
set_shape_stride_ABC(
|
||||
TensorExtent shape_act,
|
||||
TensorStride stride_act,
|
||||
TensorExtent shape_flt,
|
||||
TensorStride stride_flt,
|
||||
TensorExtent shape_xformed_act,
|
||||
TensorStride stride_xformed_act) {
|
||||
|
||||
if constexpr (ConvOp == cutlass::conv::Operator::kFprop) {
|
||||
shape_A = shape_act;
|
||||
stride_A = stride_act;
|
||||
shape_B = shape_flt;
|
||||
stride_B = stride_flt;
|
||||
shape_C = shape_xformed_act;
|
||||
stride_C = stride_xformed_act;
|
||||
}
|
||||
else if constexpr (ConvOp == cutlass::conv::Operator::kDgrad) {
|
||||
shape_A = shape_xformed_act;
|
||||
stride_A = stride_xformed_act;
|
||||
shape_B = shape_flt;
|
||||
stride_B = stride_flt;
|
||||
shape_C = shape_act;
|
||||
stride_C = stride_act;
|
||||
}
|
||||
else if constexpr (ConvOp == cutlass::conv::Operator::kWgrad) {
|
||||
shape_A = shape_xformed_act;
|
||||
stride_A = stride_xformed_act;
|
||||
shape_B = shape_act;
|
||||
stride_B = stride_act;
|
||||
shape_C = shape_flt;
|
||||
stride_C = stride_flt;
|
||||
}
|
||||
}
|
||||
|
||||
// Get problem shape MNK according to following table:
|
||||
// | | Fprop | Dgrad | Wgrad |
|
||||
// | ---- | --------- | -------- | -------- |
|
||||
// | Shape_M | (Q,P,Z,N) | (W,H,D,N) | (K) |
|
||||
// | Shape_N | (K) | (C) | (C,S,R,T) |
|
||||
// | Shape_K | (C,S,R,T) | (K,S,R,T) | (Q,P,Z,N) |
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_transformed_problem_shape_MNK() const {
|
||||
using cute::insert;
|
||||
using cute::make_shape;
|
||||
using cute::reverse;
|
||||
using cute::take;
|
||||
|
||||
if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
auto M_xformed = shape_C[0];
|
||||
auto N_xformed = reverse(take<1, RankT>(shape_C));
|
||||
auto K_xformed = reverse(take<0, RankT - 1>(shape_A));
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed);
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kFprop){
|
||||
auto M_xformed = reverse(take<0, RankT - 1>(shape_C));
|
||||
auto N_xformed = shape_C[RankT - 1];
|
||||
auto K_xformed = reverse(take<1, RankT>(shape_B));
|
||||
|
||||
return make_shape(M_xformed, N_xformed, K_xformed);
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kDgrad) {
|
||||
auto M_xformed = reverse(take<0,RankT - 1>(shape_C));
|
||||
auto N_xformed = shape_C[RankT - 1];
|
||||
// shape_B: [K,T,R,S,C], K_xformed: [K,S,R,T]
|
||||
auto K_xformed = insert<0>(
|
||||
(reverse(take<1,RankT - 1>(shape_B))),
|
||||
shape_B[0]);
|
||||
return make_shape(M_xformed, N_xformed, K_xformed);
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
// Get A extents.
|
||||
// fprop: A extents array contains [N,D,H,W,C]. Turn that into ((W,H,D,N), (C))
|
||||
// wgrad: A extents array contains [N,Z,P,Q,K]. Turn that into ((K), (Q,P,Z,N))
|
||||
// dgrad: A extents array contains [N,Z,P,Q,K]. Turn that into ((Q,P,Z,N), (K))
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_shape_A() const {
|
||||
using cute::make_shape;
|
||||
using cute::take;
|
||||
|
||||
if constexpr (ConvOp == conv::Operator::kFprop ||
|
||||
ConvOp == conv::Operator::kDgrad) {
|
||||
return make_shape(
|
||||
cute::reverse(take<0, RankT - 1>(shape_A)),
|
||||
shape_A[RankT - 1]);
|
||||
}
|
||||
// For wgrad kernel, we need to linearize NZPQ for tensor A
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
return make_shape(
|
||||
shape_A[RankT - 1],
|
||||
cute::product(take<0, RankT - 1>(shape_A)));
|
||||
}
|
||||
}
|
||||
|
||||
// Get B extents.
|
||||
// fprop: B extents array contains [K,T,R,S,C]. Turn that into ((K), (C,S,R,T))
|
||||
// wgrad: B extents array contains [N,D,H,W,C]. Turn that into ((C), (W,H,D,N))
|
||||
// dgrad: B extents array contains [K,T,R,S,C]. Turn that into ((C), (K,S,R,T))
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
get_shape_B() const {
|
||||
using cute::make_shape;
|
||||
using cute::reverse;
|
||||
using cute::take;
|
||||
|
||||
if constexpr (ConvOp == conv::Operator::kFprop) {
|
||||
return make_shape(
|
||||
shape_B[0],
|
||||
reverse(take<1, RankT>(shape_B)));
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kWgrad) {
|
||||
return make_shape(
|
||||
shape_B[RankT - 1],
|
||||
reverse(take<0, RankT - 1>(shape_B)));
|
||||
}
|
||||
else if constexpr (ConvOp == conv::Operator::kDgrad) {
|
||||
// shape_B: [K,T,R,S,C], return: [(C),(K,S,R,T)]
|
||||
return make_shape(
|
||||
shape_B[RankT - 1],
|
||||
cute::insert<0>(
|
||||
reverse(take<1, RankT - 1>(shape_B)),
|
||||
shape_B[0]));
|
||||
}
|
||||
}
|
||||
|
||||
// Static method that returns the canonical strides of tensors (layouts are right major and compact)
|
||||
CUTLASS_HOST_DEVICE
|
||||
static constexpr TensorStride
|
||||
packed_stride_right_major(TensorExtent const& extents) {
|
||||
TensorStride strides{};
|
||||
strides[RankT-1] = 1;
|
||||
cute::for_each(cute::make_rseq<RankT-1>{}, [&](auto i) {
|
||||
strides[i] = extents[i+1] * strides[i+1];
|
||||
});
|
||||
return strides;
|
||||
}
|
||||
|
||||
// Static method that returns the packed logical size of any TensorExtent
|
||||
CUTLASS_HOST_DEVICE
|
||||
static constexpr size_t
|
||||
size(TensorExtent const& extents) {
|
||||
size_t size = 1;
|
||||
cute::for_each(cute::make_seq<RankT>{}, [&](auto i) {
|
||||
size *= extents[i];
|
||||
});
|
||||
return size;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr size_t
|
||||
size_A() const {
|
||||
return shape_A[0] * stride_A[0];
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr size_t
|
||||
size_B() const {
|
||||
return shape_B[0] * stride_B[0];
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr size_t
|
||||
size_C() const {
|
||||
return shape_C[0] * stride_C[0];
|
||||
}
|
||||
|
||||
// Equality operator
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool operator==(ConvProblemShape<ConvOp, NumSpatialDimensions> const& rhs) const {
|
||||
using cute::for_each;
|
||||
using cute::make_seq;
|
||||
|
||||
bool is_equal = true;
|
||||
|
||||
// Compare all tensor extents
|
||||
for_each(make_seq<RankT>{}, [&](auto i) {
|
||||
is_equal = is_equal
|
||||
&& (shape_A[i] == rhs.shape_A[i])
|
||||
&& (shape_B[i] == rhs.shape_B[i]);
|
||||
});
|
||||
|
||||
// Compare all spatial extents
|
||||
for_each(make_seq<RankS>{}, [&](auto i) {
|
||||
is_equal = is_equal
|
||||
&& (lower_padding[i] == rhs.lower_padding[i])
|
||||
&& (upper_padding[i] == rhs.upper_padding[i])
|
||||
&& (traversal_stride[i] == rhs.traversal_stride[i])
|
||||
&& (dilation[i] == rhs.dilation[i]);
|
||||
});
|
||||
|
||||
return is_equal;
|
||||
}
|
||||
|
||||
/// Inequality operator
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool operator!=(ConvProblemShape<ConvOp, NumSpatialDimensions> const &rhs) const {
|
||||
return !(*this == rhs);
|
||||
}
|
||||
|
||||
private:
|
||||
CUTLASS_HOST_DEVICE
|
||||
constexpr auto
|
||||
calculate_xformed_act(TensorExtent shape_act, TensorExtent shape_flt) {
|
||||
TensorExtent shape_xformed_act{};
|
||||
// calculate n,z,p,q,k.
|
||||
// a helper lambda to compute a single spatial extent of the nzpqk tensor
|
||||
auto nzpqk_extent = [](int act_ext, int filter_ext, int pad_total, int dilation, int tstride) {
|
||||
return 1 + (act_ext + pad_total - ((filter_ext -1) * dilation + 1)) / tstride;
|
||||
};
|
||||
|
||||
shape_xformed_act[0] = shape_act[0]; // Activation N extent
|
||||
cute::for_each(cute::make_seq<RankS>{}, [&](auto i) {
|
||||
shape_xformed_act[i+1] = nzpqk_extent(
|
||||
shape_act[i+1], shape_flt[i+1], upper_padding[i] + lower_padding[i], dilation[i], traversal_stride[i]);
|
||||
});
|
||||
shape_xformed_act[RankT-1] = shape_flt[0]; // Filter K extent
|
||||
|
||||
TensorStride stride_xformed_act = packed_stride_right_major(shape_xformed_act);
|
||||
|
||||
return cute::make_tuple(shape_xformed_act, stride_xformed_act);
|
||||
}
|
||||
};
|
||||
|
||||
template<
|
||||
conv::Operator ConvOp,
|
||||
int SpatialDim
|
||||
>
|
||||
void print(ConvProblemShape<ConvOp, SpatialDim> const& problem) {
|
||||
printf("ConvProblemShape with %d spatial dimensions implementing cutlass::conv::Operator::%d\n",
|
||||
SpatialDim, int(ConvOp));
|
||||
printf("\tTensorA: ");
|
||||
cute::print(problem.shape_A); printf(":");
|
||||
cute::print(problem.stride_A); printf("\n");
|
||||
printf("\tTensorB: ");
|
||||
cute::print(problem.shape_B); printf(":");
|
||||
cute::print(problem.stride_B); printf("\n");
|
||||
printf("\tTensorC: ");
|
||||
cute::print(problem.shape_C); printf(":");
|
||||
cute::print(problem.stride_C); printf("\n");
|
||||
printf("\tLower padding: "); print(problem.lower_padding); printf("\n");
|
||||
printf("\tUpper padding: "); print(problem.upper_padding); printf("\n");
|
||||
printf("\tTraversal strides: "); print(problem.traversal_stride); printf("\n");
|
||||
printf("\tDilation: "); print(problem.dilation); printf("\n");
|
||||
}
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -171,6 +171,30 @@ struct TensorNHWCShape {
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Shape of a conv2d stride, which controls how the filter convolves around the input volume
|
||||
template <
|
||||
/// Stride in horizontal direction
|
||||
int u = 1,
|
||||
/// Stride in vertical direction
|
||||
int v = 1
|
||||
>
|
||||
struct Stride2D {
|
||||
static int const kU = u;
|
||||
static int const kV = v;
|
||||
|
||||
//
|
||||
// Static member functions
|
||||
//
|
||||
|
||||
/// Returns a Coord object
|
||||
CUTLASS_HOST_DEVICE
|
||||
static Coord<2> toCoord() {
|
||||
return make_Coord(kU, kV);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace conv
|
||||
|
||||
@@ -0,0 +1,414 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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
|
||||
|
||||
// common
|
||||
#include "cutlass/arch/mma.h"
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/arch/mma.h"
|
||||
#include "cutlass/trace.h"
|
||||
#include "cutlass/cluster_launch.hpp"
|
||||
#include "cutlass/device_kernel.h"
|
||||
|
||||
#include "cutlass/conv/kernel/conv_universal.hpp"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/detail/layout.hpp"
|
||||
#include "cutlass/cuda_host_adapter.hpp"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv::device {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/*!
|
||||
ConvUniversalAdapter is a stateful, reusable handle built around a kernel
|
||||
of type cutlass::conv::kernel::ConvUniversal.
|
||||
|
||||
It manages the lifetime of the underlying `kernel::Params` struct, and exposes APIs
|
||||
to create it from the host facing arguments. For power users, static methods
|
||||
are exposed that bypass the stateful methods or args->params lowering.
|
||||
*/
|
||||
template <class ConvKernel_>
|
||||
class ConvUniversalAdapter
|
||||
{
|
||||
public:
|
||||
using ConvKernel = ConvKernel_;
|
||||
using TileShape = typename ConvKernel::TileShape;
|
||||
using ElementA = typename ConvKernel::ElementA;
|
||||
using ElementB = typename ConvKernel::ElementB;
|
||||
using ElementC = typename ConvKernel::ElementC;
|
||||
using ElementD = typename ConvKernel::ElementD;
|
||||
using ElementAccumulator = typename ConvKernel::TiledMma::ValTypeC;
|
||||
using DispatchPolicy = typename ConvKernel::DispatchPolicy;
|
||||
using CollectiveMainloop = typename ConvKernel::CollectiveMainloop;
|
||||
using CollectiveEpilogue = typename ConvKernel::CollectiveEpilogue;
|
||||
|
||||
static bool const kEnableCudaHostAdapter = CUTLASS_ENABLE_CUDA_HOST_ADAPTER;
|
||||
|
||||
// Tease out meta-information about the conv algorithm
|
||||
static constexpr conv::Operator kConvolutionalOperator = DispatchPolicy::ConvOp;
|
||||
static constexpr int NumSpatialDimensions = ConvKernel::NumSpatialDimensions;
|
||||
|
||||
// If our TiledMMA's instruction thread layout size is larger than 1, we know its a tensorop!
|
||||
using OperatorClass = cute::conditional_t<
|
||||
(cute::size(typename ConvKernel::TiledMma::AtomThrID{}) > 1),
|
||||
cutlass::arch::OpClassTensorOp, cutlass::arch::OpClassSimt>;
|
||||
|
||||
using ArchTag = typename ConvKernel::ArchTag;
|
||||
|
||||
// Assume TiledMma's ShapeMNK is the same as 2.x's ThreadblockShape
|
||||
using ThreadblockShape = cutlass::gemm::GemmShape<
|
||||
cute::size<0>(TileShape{}),
|
||||
cute::size<1>(TileShape{}),
|
||||
cute::size<2>(TileShape{})>;
|
||||
|
||||
using ClusterShape = cutlass::gemm::GemmShape<
|
||||
cute::size<0>(typename ConvKernel::DispatchPolicy::ClusterShape{}),
|
||||
cute::size<1>(typename ConvKernel::DispatchPolicy::ClusterShape{}),
|
||||
cute::size<2>(typename ConvKernel::DispatchPolicy::ClusterShape{})>;
|
||||
|
||||
// Instruction shape is easy too, since we get that directly from our TiledMma's atom shape
|
||||
using InstructionShape = cutlass::gemm::GemmShape<
|
||||
cute::size<0>(typename CollectiveMainloop::TiledMma::AtomShape_MNK{}),
|
||||
cute::size<1>(typename CollectiveMainloop::TiledMma::AtomShape_MNK{}),
|
||||
cute::size<2>(typename CollectiveMainloop::TiledMma::AtomShape_MNK{})>;
|
||||
|
||||
// Legacy: provide a correct warp count, but no reliable warp shape
|
||||
static int const kThreadCount = ConvKernel::MaxThreadsPerBlock;
|
||||
|
||||
// Warp shape is not a primary API type in 3.x
|
||||
// But we can best approximate it by inspecting the TiledMma
|
||||
// For this, we make the assumption that we always have 4 warps along M, and rest along N, none along K
|
||||
// We also always round up the warp count to 4 if the tiled mma is smaller than 128 threads
|
||||
static constexpr int WarpsInMma = cute::max(4, CUTE_STATIC_V(cute::size(typename ConvKernel::TiledMma{})) / 32);
|
||||
static constexpr int WarpsInMmaM = 4;
|
||||
static constexpr int WarpsInMmaN = cute::ceil_div(WarpsInMma, WarpsInMmaM);
|
||||
using WarpCount = cutlass::gemm::GemmShape<WarpsInMmaM, WarpsInMmaN, 1>;
|
||||
using WarpShape = cutlass::gemm::GemmShape<
|
||||
CUTE_STATIC_V(cute::tile_size<0>(typename CollectiveMainloop::TiledMma{})) / WarpsInMmaM,
|
||||
CUTE_STATIC_V(cute::tile_size<1>(typename CollectiveMainloop::TiledMma{})) / WarpsInMmaN,
|
||||
CUTE_STATIC_V(cute::tile_size<2>(typename CollectiveMainloop::TiledMma{}))>;
|
||||
|
||||
static int constexpr kStages = CollectiveMainloop::DispatchPolicy::Stages;
|
||||
|
||||
// Inspect TiledCopy for A and B to compute the alignment size
|
||||
static int constexpr kAlignmentA = detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
typename CollectiveMainloop::GmemTiledCopyA, ElementA>();
|
||||
static int constexpr kAlignmentB = detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
typename CollectiveMainloop::GmemTiledCopyB, ElementB>();
|
||||
static int constexpr kAlignmentC = detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
typename CollectiveEpilogue::GmemTiledCopyC, ElementC>();
|
||||
static int constexpr kAlignmentD = detail::get_alignment_count_from_gmem_tiled_copy<
|
||||
typename CollectiveEpilogue::GmemTiledCopyD, ElementD>();
|
||||
|
||||
using EpilogueOutputOp = typename CollectiveEpilogue::ThreadEpilogueOp;
|
||||
|
||||
/// Argument structure: User API
|
||||
using Arguments = typename ConvKernel::Arguments;
|
||||
/// Argument structure: Kernel API
|
||||
using Params = typename ConvKernel::Params;
|
||||
|
||||
private:
|
||||
|
||||
/// Kernel API parameters object
|
||||
Params params_;
|
||||
|
||||
public:
|
||||
|
||||
/// Determines whether the conv can execute the given problem.
|
||||
static Status
|
||||
can_implement(Arguments const& args) {
|
||||
if (ConvKernel::can_implement(args)) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
else {
|
||||
return Status::kInvalid;
|
||||
}
|
||||
}
|
||||
|
||||
/// Gets the workspace size
|
||||
static size_t
|
||||
get_workspace_size(Arguments const& args) {
|
||||
size_t workspace_bytes = 0;
|
||||
CUTLASS_TRACE_HOST(" workspace_bytes: " << workspace_bytes);
|
||||
|
||||
workspace_bytes += ConvKernel::get_workspace_size(args);
|
||||
return workspace_bytes;
|
||||
}
|
||||
|
||||
/// Computes the grid shape
|
||||
static dim3
|
||||
get_grid_shape(Arguments const& args, void* workspace = nullptr) {
|
||||
auto tmp_params = ConvKernel::to_underlying_arguments(args, workspace);
|
||||
return ConvKernel::get_grid_shape(tmp_params);
|
||||
}
|
||||
|
||||
/// Computes the grid shape
|
||||
static dim3
|
||||
get_grid_shape(Params const& params) {
|
||||
return ConvKernel::get_grid_shape(params);
|
||||
}
|
||||
|
||||
/// Computes the maximum number of active blocks per multiprocessor
|
||||
static int maximum_active_blocks(int /* smem_capacity */ = -1) {
|
||||
CUTLASS_TRACE_HOST("ConvUniversal::maximum_active_blocks()");
|
||||
int max_active_blocks = -1;
|
||||
int smem_size = ConvKernel::SharedStorageSize;
|
||||
|
||||
// first, account for dynamic smem capacity if needed
|
||||
cudaError_t result;
|
||||
if (smem_size >= (48 << 10)) {
|
||||
CUTLASS_TRACE_HOST(" Setting smem size to " << smem_size);
|
||||
result = cudaFuncSetAttribute(
|
||||
device_kernel<ConvKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
if (cudaSuccess != result) {
|
||||
result = cudaGetLastError(); // to clear the error bit
|
||||
CUTLASS_TRACE_HOST(
|
||||
" cudaFuncSetAttribute() returned error: "
|
||||
<< cudaGetErrorString(result));
|
||||
return -1;
|
||||
}
|
||||
}
|
||||
|
||||
// query occupancy after setting smem size
|
||||
result = cudaOccupancyMaxActiveBlocksPerMultiprocessor(
|
||||
&max_active_blocks,
|
||||
device_kernel<ConvKernel>,
|
||||
ConvKernel::MaxThreadsPerBlock,
|
||||
smem_size);
|
||||
|
||||
if (cudaSuccess != result) {
|
||||
result = cudaGetLastError(); // to clear the error bit
|
||||
CUTLASS_TRACE_HOST(
|
||||
" cudaOccupancyMaxActiveBlocksPerMultiprocessor() returned error: "
|
||||
<< cudaGetErrorString(result));
|
||||
return -1;
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST(" max_active_blocks: " << max_active_blocks);
|
||||
return max_active_blocks;
|
||||
}
|
||||
|
||||
/// Initializes conv state from arguments.
|
||||
Status
|
||||
initialize(
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr,
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
CUTLASS_TRACE_HOST("ConvUniversal::initialize() - workspace "
|
||||
<< workspace << ", stream: " << (stream ? "non-null" : "null"));
|
||||
|
||||
size_t workspace_bytes = ConvKernel::get_workspace_size(args);
|
||||
CUTLASS_TRACE_HOST(" workspace_bytes: " << workspace_bytes);
|
||||
|
||||
if (workspace_bytes) {
|
||||
if (!workspace) {
|
||||
CUTLASS_TRACE_HOST(" error: device workspace must not be null");
|
||||
return Status::kErrorWorkspaceNull;
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST(" clearing device workspace");
|
||||
cudaError_t result = cudaMemsetAsync(workspace, 0, workspace_bytes, stream);
|
||||
if (cudaSuccess != result) {
|
||||
result = cudaGetLastError(); // to clear the error bit
|
||||
CUTLASS_TRACE_HOST(" cudaMemsetAsync() returned error " << cudaGetErrorString(result));
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
|
||||
// Initialize the Params structure
|
||||
params_ = ConvKernel::to_underlying_arguments(args, workspace);
|
||||
|
||||
// Don't set the function attributes - require the CudaHostAdapter to set it.
|
||||
if constexpr (kEnableCudaHostAdapter) {
|
||||
CUTLASS_ASSERT(cuda_adapter);
|
||||
return Status::kSuccess;
|
||||
}
|
||||
else {
|
||||
// account for dynamic smem capacity if needed
|
||||
int smem_size = ConvKernel::SharedStorageSize;
|
||||
if (smem_size >= (48 << 10)) {
|
||||
CUTLASS_TRACE_HOST(" Setting smem size to " << smem_size);
|
||||
cudaError_t result = cudaFuncSetAttribute(
|
||||
device_kernel<ConvKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
if (cudaSuccess != result) {
|
||||
result = cudaGetLastError(); // to clear the error bit
|
||||
CUTLASS_TRACE_HOST(" cudaFuncSetAttribute() returned error: " << cudaGetErrorString(result));
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
}
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Update API is preserved in 3.0, but does not guarantee a lightweight update of params.
|
||||
Status
|
||||
update(Arguments const& args, void* workspace = nullptr) {
|
||||
CUTLASS_TRACE_HOST("ConvUniversal()::update() - workspace: " << workspace);
|
||||
|
||||
size_t workspace_bytes = get_workspace_size(args);
|
||||
if (workspace_bytes > 0 && nullptr == workspace) {
|
||||
return Status::kErrorWorkspaceNull;
|
||||
}
|
||||
|
||||
params_ = ConvKernel::to_underlying_arguments(args, workspace);
|
||||
return Status::kSuccess;
|
||||
}
|
||||
|
||||
/// Primary run() entry point API that is static allowing users to create and manage their own params.
|
||||
/// Supplied params struct must be construct by calling ConvKernel::to_underling_arguments()
|
||||
static Status
|
||||
run(Params& params, cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
CUTLASS_TRACE_HOST("ConvUniversal::run()");
|
||||
dim3 const block = ConvKernel::get_block_shape();
|
||||
dim3 const grid = get_grid_shape(params);
|
||||
|
||||
// configure smem size and carveout
|
||||
int smem_size = ConvKernel::SharedStorageSize;
|
||||
|
||||
Status launch_result;
|
||||
// Use extended launch API only for mainloops that use it
|
||||
if constexpr(ConvKernel::ArchTag::kMinComputeCapability >= 90) {
|
||||
dim3 cluster(cute::size<0>(typename ConvKernel::DispatchPolicy::ClusterShape{}),
|
||||
cute::size<1>(typename ConvKernel::DispatchPolicy::ClusterShape{}),
|
||||
cute::size<2>(typename ConvKernel::DispatchPolicy::ClusterShape{}));
|
||||
void* kernel_params[] = {¶ms};
|
||||
if constexpr (kEnableCudaHostAdapter) {
|
||||
//
|
||||
// Use the cuda host adapter
|
||||
//
|
||||
CUTLASS_ASSERT(cuda_adapter);
|
||||
if (cuda_adapter) {
|
||||
|
||||
launch_result = cuda_adapter->launch(
|
||||
grid, cluster, block, smem_size, stream, kernel_params, 0
|
||||
);
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
else {
|
||||
|
||||
CUTLASS_ASSERT(cuda_adapter == nullptr);
|
||||
void const* kernel = (void const*) device_kernel<ConvKernel>;
|
||||
|
||||
launch_result = ClusterLauncher::launch(
|
||||
grid, cluster, block, smem_size, stream, kernel, kernel_params);
|
||||
|
||||
}
|
||||
}
|
||||
else {
|
||||
launch_result = Status::kSuccess;
|
||||
|
||||
if constexpr (kEnableCudaHostAdapter) {
|
||||
CUTLASS_ASSERT(cuda_adapter);
|
||||
if (cuda_adapter) {
|
||||
void* kernel_params[] = {¶ms};
|
||||
|
||||
launch_result = cuda_adapter->launch(
|
||||
grid, block, smem_size, stream, kernel_params, 0
|
||||
);
|
||||
|
||||
}
|
||||
else {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
else {
|
||||
CUTLASS_ASSERT(cuda_adapter == nullptr);
|
||||
device_kernel<ConvKernel><<<grid, block, smem_size, stream>>>(params);
|
||||
}
|
||||
}
|
||||
|
||||
cudaError_t result = cudaGetLastError();
|
||||
if (cudaSuccess == result && Status::kSuccess == launch_result) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
else {
|
||||
CUTLASS_TRACE_HOST(" Kernel launch failed. Reason: " << result);
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
|
||||
//
|
||||
// Non-static launch overloads that first create and set the internal params struct of this kernel handle.
|
||||
//
|
||||
|
||||
/// Launches the kernel after first constructing Params internal state from supplied arguments.
|
||||
Status
|
||||
run(
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr,
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr
|
||||
) {
|
||||
Status status = initialize(args, workspace, stream, cuda_adapter);
|
||||
if (Status::kSuccess == status) {
|
||||
status = run(params_, stream, cuda_adapter);
|
||||
}
|
||||
return status;
|
||||
}
|
||||
|
||||
/// Launches the kernel after first constructing Params internal state from supplied arguments.
|
||||
Status
|
||||
operator()(
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr,
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
return run(args, workspace, stream, cuda_adapter);
|
||||
}
|
||||
|
||||
/// Overload that allows a user to re-launch the same kernel without updating internal params struct.
|
||||
Status
|
||||
run(cudaStream_t stream = nullptr) {
|
||||
return run(params_, stream);
|
||||
}
|
||||
|
||||
/// Overload that allows a user to re-launch the same kernel without updating internal params struct.
|
||||
Status
|
||||
operator()(cudaStream_t stream = nullptr) {
|
||||
return run(params_, stream);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::device
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -39,6 +39,7 @@
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/device_kernel.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/cuda_host_adapter.hpp"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -80,6 +81,8 @@ public:
|
||||
static cutlass::conv::StrideSupport const kStrideSupport = UnderlyingKernel::kStrideSupport;
|
||||
static cutlass::conv::GroupMode const kGroupMode = UnderlyingKernel::kGroupMode;
|
||||
|
||||
static bool const kEnableCudaHostAdapter = CUTLASS_ENABLE_CUDA_HOST_ADAPTER;
|
||||
|
||||
static int const kWarpCount =
|
||||
(ThreadblockShape::kM / WarpShape::kM) *
|
||||
(ThreadblockShape::kN / WarpShape::kN) *
|
||||
@@ -230,7 +233,8 @@ public:
|
||||
Status initialize(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
cudaStream_t stream = nullptr) {
|
||||
cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
|
||||
if (args.problem_size.split_k_slices > 1) {
|
||||
|
||||
@@ -250,16 +254,22 @@ public:
|
||||
args,
|
||||
static_cast<int *>(workspace)
|
||||
);
|
||||
|
||||
int smem_size = int(sizeof(typename UnderlyingKernel::SharedStorage));
|
||||
|
||||
if (smem_size >= (48 << 10)) {
|
||||
cudaError_t result = cudaFuncSetAttribute(cutlass::Kernel<UnderlyingKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
if constexpr (kEnableCudaHostAdapter) {
|
||||
CUTLASS_ASSERT(cuda_adapter);
|
||||
return Status::kSuccess;
|
||||
}
|
||||
else {
|
||||
int smem_size = int(sizeof(typename UnderlyingKernel::SharedStorage));
|
||||
|
||||
if (smem_size >= (48 << 10)) {
|
||||
cudaError_t result = cudaFuncSetAttribute(cutlass::Kernel<UnderlyingKernel>,
|
||||
cudaFuncAttributeMaxDynamicSharedMemorySize,
|
||||
smem_size);
|
||||
|
||||
if (result != cudaSuccess) {
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -281,7 +291,7 @@ public:
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status run(cudaStream_t stream = nullptr) {
|
||||
Status run(cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
|
||||
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
@@ -290,29 +300,53 @@ public:
|
||||
dim3 block(32 * kWarpCount, 1, 1);
|
||||
|
||||
int smem_size = int(sizeof(typename UnderlyingKernel::SharedStorage));
|
||||
cutlass::Status launch_result = cutlass::Status::kSuccess ;
|
||||
|
||||
cutlass::Kernel<UnderlyingKernel><<<grid, block, smem_size, stream>>>(params_);
|
||||
if constexpr (kEnableCudaHostAdapter) {
|
||||
//
|
||||
// Use the cuda host adapter
|
||||
//
|
||||
CUTLASS_ASSERT(cuda_adapter);
|
||||
if (cuda_adapter) {
|
||||
|
||||
void* kernel_params[] = {¶ms_};
|
||||
launch_result = cuda_adapter->launch(
|
||||
grid, dim3(1,1,1), block, smem_size, stream, kernel_params, 0
|
||||
);
|
||||
}
|
||||
else {
|
||||
launch_result = Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
else {
|
||||
cutlass::Kernel<UnderlyingKernel><<<grid, block, smem_size, stream>>>(params_);
|
||||
}
|
||||
|
||||
cudaError_t result = cudaGetLastError();
|
||||
|
||||
return result == cudaSuccess ? Status::kSuccess : Status::kErrorInternal;
|
||||
if (cudaSuccess == result && Status::kSuccess == launch_result) {
|
||||
return Status::kSuccess;
|
||||
}
|
||||
else {
|
||||
CUTLASS_TRACE_HOST(" Kernel launch failed. Reason: " << result);
|
||||
return Status::kErrorInternal;
|
||||
}
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status operator()(cudaStream_t stream = nullptr) {
|
||||
return run(stream);
|
||||
Status operator()(cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
return run(stream, cuda_adapter);
|
||||
}
|
||||
|
||||
/// Runs the kernel using initialized state.
|
||||
Status operator()(
|
||||
Arguments const &args,
|
||||
void *workspace = nullptr,
|
||||
cudaStream_t stream = nullptr) {
|
||||
cudaStream_t stream = nullptr, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
|
||||
Status status = initialize(args, workspace, stream);
|
||||
Status status = initialize(args, workspace, stream, cuda_adapter);
|
||||
|
||||
if (status == Status::kSuccess) {
|
||||
status = run(stream);
|
||||
status = run(stream, cuda_adapter);
|
||||
}
|
||||
|
||||
return status;
|
||||
|
||||
@@ -0,0 +1,83 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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/convolution.h"
|
||||
#include "cutlass/arch/arch.h"
|
||||
|
||||
#include "cute/layout.hpp"
|
||||
#include "cute/numeric/integral_constant.hpp"
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv {
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
//
|
||||
// Policies for categorical dispatch of mainloop against kernel grid schedules
|
||||
//
|
||||
struct KernelImplicitTmaWarpSpecializedSm90 { };
|
||||
struct KernelImplicitTmaWarpSpecializedSm90Cooperative { };
|
||||
struct KernelImplicitTmaWarpSpecializedSm90Pingpong { };
|
||||
|
||||
//
|
||||
// Collective Mainloop Policies
|
||||
//
|
||||
|
||||
// n-buffer in smem (Hopper TMA), pipelined with Hopper GMMA and TMA, static schedule between TMA and GMMA
|
||||
// for fprop
|
||||
template<
|
||||
conv::Operator ConvOp_,
|
||||
int Stages_,
|
||||
int NumSpatialDimensions_,
|
||||
class ClusterShape_ = cute::Shape<cute::C<1>,cute::C<1>,cute::C<1>>,
|
||||
class KernelSchedule = KernelImplicitTmaWarpSpecializedSm90,
|
||||
int PipelineAsyncMmaStages_ = 1
|
||||
>
|
||||
struct MainloopSm90TmaGmmaWarpSpecializedImplicitGemm {
|
||||
static constexpr int Stages = Stages_;
|
||||
static constexpr int NumSpatialDimensions = NumSpatialDimensions_;
|
||||
static constexpr Operator ConvOp = ConvOp_;
|
||||
static constexpr int PipelineAsyncMmaStages = PipelineAsyncMmaStages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using ArchTag = arch::Sm90;
|
||||
using Schedule = KernelSchedule;
|
||||
|
||||
static_assert(NumSpatialDimensions >= 1);
|
||||
static_assert(! (cute::is_same_v<KernelSchedule,KernelImplicitTmaWarpSpecializedSm90Cooperative> ||
|
||||
cute::is_same_v<KernelSchedule,KernelImplicitTmaWarpSpecializedSm90Pingpong>),
|
||||
"Persistent schedules not support for conv yet.");
|
||||
};
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv
|
||||
@@ -0,0 +1,63 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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/detail/dependent_false.hpp"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv::kernel {
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/*
|
||||
* Stateless universal device CONV kernel type that treats CONV as
|
||||
* a composition of a collective mainloop and a collective epilogue.
|
||||
**/
|
||||
template <
|
||||
class CollectiveMainloop_,
|
||||
class CollectiveEpilogue_,
|
||||
class TileSchedulerTag_ = void,
|
||||
class Enable = void
|
||||
>
|
||||
class ConvUniversal {
|
||||
static_assert(cutlass::detail::dependent_false<Enable>,
|
||||
"Could not find a valid specialization at the kernel layer to dispatch against.");
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::kernel
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include "cutlass/conv/kernel/sm90_implicit_gemm_tma_warpspecialized.hpp"
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -106,6 +106,27 @@ struct DefaultConvEpilogue<
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
template <
|
||||
typename ArchTag,
|
||||
typename Shape,
|
||||
typename WarpMmaSimt,
|
||||
typename ElementOutput,
|
||||
typename ElementTensor,
|
||||
typename ElementVector,
|
||||
typename OutputOp,
|
||||
int ElementsPerAccess
|
||||
>
|
||||
struct DefaultConvEpilogueWithBroadcastSimt {
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueWithBroadcastSimt<
|
||||
Shape,
|
||||
WarpMmaSimt,
|
||||
ElementOutput,
|
||||
ElementTensor,
|
||||
ElementVector,
|
||||
OutputOp,
|
||||
ElementsPerAccess
|
||||
>::Epilogue;
|
||||
};
|
||||
|
||||
template <
|
||||
typename ArchTag,
|
||||
|
||||
@@ -0,0 +1,127 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief Defines a default configuration for convolution with absolute maximum calculation.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/conv/kernel/default_conv2d_fprop.h"
|
||||
#include "cutlass/conv/kernel/implicit_gemm_convolution_with_absmax.h"
|
||||
|
||||
#include "cutlass/epilogue/threadblock/default_epilogue_with_absmax.h"
|
||||
#include "cutlass/epilogue/threadblock/epilogue_with_absmax.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename OperatorClass,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag,
|
||||
conv::IteratorAlgorithm IteratorAlgorithm = IteratorAlgorithm::kOptimized,
|
||||
conv::StrideSupport StrideSupport = StrideSupport::kStrided,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int AlignmentA = 128 / cutlass::sizeof_bits<ElementA>::value,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int AlignmentB = 128 / cutlass::sizeof_bits<ElementB>::value
|
||||
>
|
||||
struct DefaultConv2dFpropWithAbsMax {
|
||||
|
||||
using ImplicitGemmBase = typename DefaultConv2dFprop<
|
||||
ElementA, LayoutA,
|
||||
ElementB, LayoutB,
|
||||
ElementC, LayoutC,
|
||||
ElementAccumulator,
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm,
|
||||
StrideSupport,
|
||||
AlignmentA,
|
||||
AlignmentB
|
||||
>::Kernel;
|
||||
|
||||
// Define epilogue
|
||||
using Epilogue = typename cutlass::epilogue::threadblock::DefaultEpilogueWithAbsMax<
|
||||
typename ImplicitGemmBase::Epilogue::Shape,
|
||||
typename ImplicitGemmBase::Epilogue::WarpMmaOperator,
|
||||
ImplicitGemmBase::Epilogue::kPartitionsK,
|
||||
ElementC,
|
||||
typename EpilogueOutputOp::ElementAuxOutput,
|
||||
ElementC,
|
||||
EpilogueOutputOp,
|
||||
ImplicitGemmBase::Epilogue::kElementsPerAccess
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolutionWithAbsMax<
|
||||
typename ImplicitGemmBase::Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -121,6 +121,96 @@ struct DefaultConv2dFpropWithBroadcast {
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// OpClassSimt convolutions
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Defines a kernel for Conv2dFprop specialization for Analytic IteratorAlgorithm,
|
||||
/// multi-stage pipeline, and FFMA-based mainloop for SM80
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag,
|
||||
conv::StrideSupport StrideSupport,
|
||||
int AlignmentA,
|
||||
int AlignmentB
|
||||
>
|
||||
struct DefaultConv2dFpropWithBroadcast <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kAnalytic,
|
||||
StrideSupport,
|
||||
AlignmentA,
|
||||
AlignmentB
|
||||
> {
|
||||
|
||||
using ImplicitGemmBase = typename DefaultConv2dFprop<
|
||||
ElementA, LayoutA,
|
||||
ElementB, LayoutB,
|
||||
ElementC, LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kAnalytic,
|
||||
StrideSupport,
|
||||
AlignmentA,
|
||||
AlignmentB
|
||||
>::Kernel;
|
||||
|
||||
// Define epilogue
|
||||
using Epilogue = typename cutlass::conv::kernel::detail::DefaultConvEpilogueWithBroadcastSimt<
|
||||
ArchTag,
|
||||
typename ImplicitGemmBase::Epilogue::Shape,
|
||||
typename ImplicitGemmBase::Epilogue::WarpMmaOperator,
|
||||
ElementC,
|
||||
typename EpilogueOutputOp::ElementT,
|
||||
typename EpilogueOutputOp::ElementVector,
|
||||
EpilogueOutputOp,
|
||||
ImplicitGemmBase::Epilogue::kElementsPerAccess
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolutionWithFusedEpilogue<
|
||||
typename ImplicitGemmBase::Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
|
||||
@@ -57,7 +57,7 @@ namespace conv {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Defines a kernel for Conv2dGroupFpro
|
||||
/// Defines a kernel for Conv2dGroupFprop
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
@@ -135,11 +135,11 @@ struct DefaultConv2dGroupFprop <
|
||||
AlignmentB
|
||||
> {
|
||||
|
||||
static_assert(std::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
|
||||
// Define the core components from GEMM
|
||||
@@ -269,11 +269,11 @@ struct DefaultConv2dGroupFprop <
|
||||
AlignmentB
|
||||
> {
|
||||
|
||||
static_assert(std::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
|
||||
// Define the core components from GEMM
|
||||
@@ -400,11 +400,11 @@ struct DefaultConv2dGroupFprop <
|
||||
AlignmentB
|
||||
> {
|
||||
|
||||
static_assert(std::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
|
||||
// Define the core components from GEMM
|
||||
@@ -530,11 +530,11 @@ struct DefaultConv2dGroupFprop <
|
||||
AlignmentB
|
||||
> {
|
||||
|
||||
static_assert(std::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutA, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutB, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
static_assert(std::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
static_assert(platform::is_same<LayoutC, cutlass::layout::TensorNHWC>::value,
|
||||
"Current group conv only support NHWC layout");
|
||||
|
||||
// Define the core components from GEMM
|
||||
|
||||
@@ -293,6 +293,439 @@ struct DefaultConv3dDgrad <
|
||||
};
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// OpClassSimt convolutions
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dDgrad <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kAnalytic,
|
||||
conv::StrideSupport::kStrided
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
ElementB, layout::RowMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
Stages, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::Conv3dDgradOutputGradientTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
ThreadMapA,
|
||||
conv::StrideSupport::kStrided
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::Conv3dDgradFilterTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
ThreadMapB
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmMultistage<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
arch::CacheOperation::Always,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
arch::CacheOperation::Always,
|
||||
MmaPolicy,
|
||||
Stages
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kDgrad,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv3dDgrad specialization for Optimized IteratorAlgorithm,
|
||||
/// multi-stage pipeline, and FFMA-based mainloop for SM80
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dDgrad <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kOptimized,
|
||||
StrideSupport::kUnity
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
ElementB, layout::RowMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
Stages, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::Conv3dDgradOutputGradientTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
ThreadMapA,
|
||||
StrideSupport::kUnity
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::Conv3dDgradFilterTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
ThreadMapB
|
||||
// ThreadMapB,
|
||||
// StrideSupport::kUnity
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmMultistage<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
arch::CacheOperation::Always,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
arch::CacheOperation::Always,
|
||||
MmaPolicy,
|
||||
Stages
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kDgrad,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dDgrad <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
2,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kAnalytic,
|
||||
conv::StrideSupport::kStrided
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
ElementB, layout::RowMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
2, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
// cutlass::conv::threadblock::TileIteratorStridedDgrad<
|
||||
cutlass::conv::threadblock::Conv3dDgradOutputGradientTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
ThreadMapA,
|
||||
conv::StrideSupport::kStrided
|
||||
// >
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
// cutlass::conv::threadblock::TileIteratorStridedDgrad<
|
||||
cutlass::conv::threadblock::Conv3dDgradFilterTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
ThreadMapB
|
||||
// >
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmPipelined<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
MmaPolicy
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kDgrad,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv3dDgrad specialization for Optimized IteratorAlgorithm,
|
||||
/// 2 stage pipeline, and FFMA-based mainloop for SM50
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dDgrad <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
2,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kOptimized,
|
||||
StrideSupport::kUnity
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
ElementB, layout::RowMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
2, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
// cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv3dDgradOutputGradientTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
ThreadMapA,
|
||||
StrideSupport::kUnity
|
||||
// >
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
// cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv3dDgradFilterTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
ThreadMapB
|
||||
// ThreadMapB,
|
||||
// StrideSupport::kUnity
|
||||
// >
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmPipelined<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
MmaPolicy
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kDgrad,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
|
||||
@@ -54,7 +54,7 @@ namespace conv {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Defines a kernel for Conv2dFprop
|
||||
/// Defines a kernel for Conv3dFprop
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
@@ -185,7 +185,7 @@ struct DefaultConv3dFprop <
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv2dFprop specialization for Analytic IteratorAlgorithm and multistage
|
||||
/// Defines a kernel for Conv3dFprop specialization for Analytic IteratorAlgorithm and multistage
|
||||
// pipeline.
|
||||
template <
|
||||
typename ElementA,
|
||||
@@ -506,7 +506,437 @@ struct DefaultConv3dFprop <
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// OpClassSimt convolutions
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Defines a kernel for Conv3dFprop specialization for Analytic IteratorAlgorithm,
|
||||
/// multi-stage pipeline, and FFMA-based mainloop for SM80
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dFprop <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kAnalytic
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
ElementB, layout::ColumnMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
Stages, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::Conv3dFpropActivationTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
ThreadMapA
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::Conv3dFpropFilterTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
ThreadMapB
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmMultistage<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
arch::CacheOperation::Always,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
arch::CacheOperation::Always,
|
||||
MmaPolicy,
|
||||
Stages
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv3dFprop specialization for Optimized IteratorAlgorithm,
|
||||
/// multi-stage pipeline, and FFMA-based mainloop for SM80
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dFprop <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kOptimized
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
ElementB, layout::ColumnMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
Stages, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::Conv3dFpropActivationTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ThreadMapA
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::Conv3dFpropFilterTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ThreadMapB
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmMultistage<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
arch::CacheOperation::Always,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
arch::CacheOperation::Always,
|
||||
MmaPolicy,
|
||||
Stages
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv3dFprop specialization for Analytic IteratorAlgorithm,
|
||||
/// 2 stage pipeline, and FFMA-based mainloop for SM50
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dFprop <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
2,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kAnalytic
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
ElementB, layout::ColumnMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
2, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv3dFpropActivationTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
ThreadMapA
|
||||
>
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv3dFpropFilterTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
ThreadMapB
|
||||
>
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmPipelined<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
MmaPolicy
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv3dFprop specialization for Optimized IteratorAlgorithm,
|
||||
/// 2 stage pipeline, and FFMA-based mainloop for SM50
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dFprop <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
2,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kOptimized
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::RowMajor,
|
||||
ElementB, layout::ColumnMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
2, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv3dFpropActivationTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ThreadMapA
|
||||
>
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv3dFpropFilterTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ThreadMapB
|
||||
>
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmPipelined<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
MmaPolicy
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
} // namespace kernel
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
@@ -0,0 +1,218 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief
|
||||
Defines a GEMM with Reduction based on an existing UniversalGemm kernel.
|
||||
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/conv/kernel/default_conv3d_fprop.h"
|
||||
#include "cutlass/conv/kernel/implicit_gemm_convolution_with_fused_epilogue.h"
|
||||
|
||||
#include "cutlass/epilogue/threadblock/default_epilogue_with_broadcast.h"
|
||||
#include "cutlass/epilogue/threadblock/epilogue_with_broadcast.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename OperatorClass,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag,
|
||||
conv::IteratorAlgorithm IteratorAlgorithm = IteratorAlgorithm::kOptimized,
|
||||
conv::StrideSupport StrideSupport = StrideSupport::kStrided,
|
||||
/// Access granularity of A matrix in units of elements
|
||||
int AlignmentA = 128 / cutlass::sizeof_bits<ElementA>::value,
|
||||
/// Access granularity of B matrix in units of elements
|
||||
int AlignmentB = 128 / cutlass::sizeof_bits<ElementB>::value
|
||||
>
|
||||
struct DefaultConv3dFpropWithBroadcast {
|
||||
|
||||
using ImplicitGemmBase = typename DefaultConv3dFprop<
|
||||
ElementA, LayoutA,
|
||||
ElementB, LayoutB,
|
||||
ElementC, LayoutC,
|
||||
ElementAccumulator,
|
||||
OperatorClass,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm,
|
||||
StrideSupport
|
||||
>::Kernel;
|
||||
|
||||
// Define epilogue
|
||||
using Epilogue = typename cutlass::conv::kernel::detail::DefaultConvEpilogueWithBroadcastTensorOp<
|
||||
ArchTag,
|
||||
typename ImplicitGemmBase::Epilogue::Shape,
|
||||
typename ImplicitGemmBase::Epilogue::WarpMmaOperator,
|
||||
ImplicitGemmBase::Epilogue::kPartitionsK,
|
||||
ElementC,
|
||||
typename EpilogueOutputOp::ElementT,
|
||||
typename EpilogueOutputOp::ElementVector,
|
||||
EpilogueOutputOp,
|
||||
ImplicitGemmBase::Epilogue::kElementsPerAccess
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolutionWithFusedEpilogue<
|
||||
typename ImplicitGemmBase::Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// OpClassSimt convolutions
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Defines a kernel for Conv3dFprop specialization for Analytic IteratorAlgorithm,
|
||||
/// multi-stage pipeline, and FFMA-based mainloop for SM80
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag,
|
||||
conv::StrideSupport StrideSupport,
|
||||
int AlignmentA,
|
||||
int AlignmentB
|
||||
>
|
||||
struct DefaultConv3dFpropWithBroadcast <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kAnalytic,
|
||||
StrideSupport,
|
||||
AlignmentA,
|
||||
AlignmentB
|
||||
> {
|
||||
|
||||
using ImplicitGemmBase = typename DefaultConv3dFprop<
|
||||
ElementA, LayoutA,
|
||||
ElementB, LayoutB,
|
||||
ElementC, LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kAnalytic,
|
||||
StrideSupport
|
||||
>::Kernel;
|
||||
|
||||
// Define epilogue
|
||||
using Epilogue = typename cutlass::conv::kernel::detail::DefaultConvEpilogueWithBroadcastSimt<
|
||||
ArchTag,
|
||||
typename ImplicitGemmBase::Epilogue::Shape,
|
||||
typename ImplicitGemmBase::Epilogue::WarpMmaOperator,
|
||||
ElementC,
|
||||
typename EpilogueOutputOp::ElementT,
|
||||
typename EpilogueOutputOp::ElementVector,
|
||||
EpilogueOutputOp,
|
||||
ImplicitGemmBase::Epilogue::kElementsPerAccess
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolutionWithFusedEpilogue<
|
||||
typename ImplicitGemmBase::Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kFprop,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -53,7 +53,7 @@ namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv2dWgrad
|
||||
/// Defines a kernel for Conv3dWgrad
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
@@ -500,6 +500,433 @@ struct DefaultConv3dWgrad <
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
};
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// OpClassSimt convolutions
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
/// Defines a kernel for Conv3dWgrad specialization for Analytic IteratorAlgorithm,
|
||||
/// multi-stage pipeline, and FFMA-based mainloop for SM80
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dWgrad <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kAnalytic
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::ColumnMajor,
|
||||
ElementB, layout::RowMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
Stages, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::Conv3dWgradOutputGradientTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
ThreadMapA
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::Conv3dWgradActivationTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
ThreadMapB
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmMultistage<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
arch::CacheOperation::Always,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
arch::CacheOperation::Always,
|
||||
MmaPolicy,
|
||||
Stages
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kWgrad,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv3dWgrad specialization for Optimized IteratorAlgorithm,
|
||||
/// multi-stage pipeline, and FFMA-based mainloop for SM80
|
||||
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
int Stages,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dWgrad <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
Stages,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kOptimized
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::ColumnMajor,
|
||||
ElementB, layout::RowMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
Stages, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::Conv3dWgradOutputGradientTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
ThreadMapA
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::Conv3dWgradActivationTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
ThreadMapB
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmMultistage<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
arch::CacheOperation::Always,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
arch::CacheOperation::Always,
|
||||
MmaPolicy,
|
||||
Stages
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kWgrad,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv3dWgrad specialization for Analytic IteratorAlgorithm,
|
||||
/// 2 stage pipeline, and FFMA-based mainloop for SM50
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dWgrad <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
2,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kAnalytic
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::ColumnMajor,
|
||||
ElementB, layout::RowMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
2, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv3dWgradOutputGradientTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
ThreadMapA
|
||||
>
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv3dWgradActivationTileAccessIteratorAnalytic<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
ThreadMapB
|
||||
>
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmPipelined<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
MmaPolicy
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kWgrad,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
/// Defines a kernel for Conv3dWgrad specialization for Optimized IteratorAlgorithm,
|
||||
/// 2 stage pipeline, and FFMA-based mainloop for SM50
|
||||
template <
|
||||
typename ElementA,
|
||||
typename LayoutA,
|
||||
typename ElementB,
|
||||
typename LayoutB,
|
||||
typename ElementC,
|
||||
typename LayoutC,
|
||||
typename ElementAccumulator,
|
||||
typename ArchTag,
|
||||
typename ThreadblockShape,
|
||||
typename WarpShape,
|
||||
typename InstructionShape,
|
||||
typename EpilogueOutputOp,
|
||||
typename ThreadblockSwizzle,
|
||||
typename MathOperatorTag
|
||||
>
|
||||
struct DefaultConv3dWgrad <
|
||||
ElementA,
|
||||
LayoutA,
|
||||
ElementB,
|
||||
LayoutB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
ElementAccumulator,
|
||||
arch::OpClassSimt,
|
||||
ArchTag,
|
||||
ThreadblockShape,
|
||||
WarpShape,
|
||||
InstructionShape,
|
||||
EpilogueOutputOp,
|
||||
ThreadblockSwizzle,
|
||||
2,
|
||||
MathOperatorTag,
|
||||
IteratorAlgorithm::kOptimized
|
||||
> {
|
||||
|
||||
// Define the core components from GEMM
|
||||
using MmaCore = typename cutlass::gemm::threadblock::DefaultMmaCore<
|
||||
ThreadblockShape, WarpShape, InstructionShape, ElementA, layout::ColumnMajor,
|
||||
ElementB, layout::RowMajor, ElementAccumulator, layout::RowMajor, arch::OpClassSimt,
|
||||
2, MathOperatorTag>;
|
||||
|
||||
// Define iterators over tiles from the A operand
|
||||
using ThreadMapA = typename MmaCore::IteratorThreadMapA;
|
||||
using IteratorA =
|
||||
cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv3dWgradOutputGradientTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kM, ThreadblockShape::kK>,
|
||||
ElementA,
|
||||
ThreadMapA
|
||||
>
|
||||
>;
|
||||
|
||||
using SmemIteratorA = typename MmaCore::SmemIteratorA;
|
||||
|
||||
// Define iterators over tiles from the B operand
|
||||
using ThreadMapB = typename MmaCore::IteratorThreadMapB;
|
||||
using IteratorB =
|
||||
cutlass::conv::threadblock::TileIterator<
|
||||
cutlass::conv::threadblock::Conv3dWgradActivationTileAccessIteratorOptimized<
|
||||
cutlass::MatrixShape<ThreadblockShape::kK, ThreadblockShape::kN>,
|
||||
ElementB,
|
||||
ThreadMapB
|
||||
>
|
||||
>;
|
||||
|
||||
using SmemIteratorB = typename MmaCore::SmemIteratorB;
|
||||
|
||||
// Warp-level GEMM components
|
||||
using WarpMmaSimtOp = typename MmaCore::MmaWarpSimt;
|
||||
using MmaPolicy = typename MmaCore::MmaPolicy;
|
||||
|
||||
// Define the Mma
|
||||
using Mma = threadblock::ImplicitGemmPipelined<
|
||||
ThreadblockShape,
|
||||
IteratorA,
|
||||
SmemIteratorA,
|
||||
IteratorB,
|
||||
SmemIteratorB,
|
||||
ElementC,
|
||||
LayoutC,
|
||||
MmaPolicy
|
||||
>;
|
||||
|
||||
// Define the epilogue
|
||||
using Epilogue = typename epilogue::threadblock::DefaultEpilogueSimt<
|
||||
ThreadblockShape,
|
||||
WarpMmaSimtOp,
|
||||
EpilogueOutputOp,
|
||||
EpilogueOutputOp::kCount
|
||||
>::Epilogue;
|
||||
|
||||
// Define the kernel
|
||||
using Kernel = cutlass::conv::kernel::ImplicitGemmConvolution<
|
||||
Mma,
|
||||
Epilogue,
|
||||
ThreadblockSwizzle,
|
||||
conv::Operator::kWgrad,
|
||||
Conv3dProblemSize
|
||||
>;
|
||||
|
||||
};
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
|
||||
@@ -202,32 +202,30 @@ struct ImplicitGemmConvolutionFusion {
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
ConvProblemSize problem_size;
|
||||
cutlass::gemm::GemmCoord grid_tiled_shape;
|
||||
gemm::GemmCoord implicit_gemm_problem_size;
|
||||
int swizzle_log_tile;
|
||||
int gemm_k_iterations;
|
||||
typename Mma::IteratorA::Params iterator_A;
|
||||
typename Mma::IteratorA::Element const *ptr_A;
|
||||
typename Mma::IteratorB::Params iterator_B;
|
||||
typename Mma::IteratorB::Element const *ptr_B;
|
||||
typename Mma::IteratorScaleBias::Params iterator_scale_bias;
|
||||
typename Mma::IteratorScaleBias::Element const *ptr_scale;
|
||||
typename Mma::IteratorScaleBias::Element const *ptr_bias;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_C;
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_C;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_D;
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_D;
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
int *semaphore;
|
||||
SplitKMode split_k_mode;
|
||||
ConvProblemSize problem_size{};
|
||||
cutlass::gemm::GemmCoord grid_tiled_shape{};
|
||||
gemm::GemmCoord implicit_gemm_problem_size{};
|
||||
int swizzle_log_tile{0};
|
||||
int gemm_k_iterations{0};
|
||||
typename Mma::IteratorA::Params iterator_A{};
|
||||
typename Mma::IteratorA::Element const *ptr_A = nullptr;
|
||||
typename Mma::IteratorB::Params iterator_B{};
|
||||
typename Mma::IteratorB::Element const *ptr_B = nullptr;
|
||||
typename Mma::IteratorScaleBias::Params iterator_scale_bias{};
|
||||
typename Mma::IteratorScaleBias::Element const *ptr_scale = nullptr;
|
||||
typename Mma::IteratorScaleBias::Element const *ptr_bias = nullptr;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_C {};
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_C = nullptr;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_D {};
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_D = nullptr;
|
||||
typename EpilogueOutputOp::Params output_op {};
|
||||
int *semaphore = nullptr;
|
||||
SplitKMode split_k_mode {};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(): swizzle_log_tile(0), gemm_k_iterations(0) { }
|
||||
Params() = default;
|
||||
|
||||
///
|
||||
CUTLASS_HOST_DEVICE
|
||||
|
||||
@@ -158,21 +158,20 @@ struct ImplicitGemmConvolutionStridedDgrad {
|
||||
// Data members
|
||||
//
|
||||
|
||||
ConvProblemSize problem_size;
|
||||
TensorRefA ref_A;
|
||||
TensorRefB ref_B;
|
||||
TensorRefC ref_C;
|
||||
TensorRefC ref_D;
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
SplitKMode split_k_mode;
|
||||
ConvProblemSize problem_size{};
|
||||
TensorRefA ref_A{};
|
||||
TensorRefB ref_B{};
|
||||
TensorRefC ref_C{};
|
||||
TensorRefC ref_D{};
|
||||
typename EpilogueOutputOp::Params output_op{};
|
||||
SplitKMode split_k_mode{};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments() { }
|
||||
Arguments() = default;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments(
|
||||
@@ -205,30 +204,28 @@ struct ImplicitGemmConvolutionStridedDgrad {
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
ConvProblemSize problem_size;
|
||||
cutlass::gemm::GemmCoord grid_tiled_shape;
|
||||
int swizzle_log_tile;
|
||||
FastDivmod stride_h_divmod;
|
||||
FastDivmod stride_w_divmod;
|
||||
int gemm_k_iterations;
|
||||
typename Mma::IteratorA::Params iterator_A;
|
||||
typename Mma::IteratorA::Element const *ptr_A;
|
||||
typename Mma::IteratorB::Params iterator_B;
|
||||
typename Mma::IteratorB::Element const *ptr_B;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_C;
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_C;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_D;
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_D;
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
int *semaphore;
|
||||
SplitKMode split_k_mode;
|
||||
ConvProblemSize problem_size{};
|
||||
cutlass::gemm::GemmCoord grid_tiled_shape{};
|
||||
int swizzle_log_tile{0};
|
||||
FastDivmod stride_h_divmod{};
|
||||
FastDivmod stride_w_divmod{};
|
||||
int gemm_k_iterations{0};
|
||||
typename Mma::IteratorA::Params iterator_A{};
|
||||
typename Mma::IteratorA::Element const *ptr_A = nullptr;
|
||||
typename Mma::IteratorB::Params iterator_B{};
|
||||
typename Mma::IteratorB::Element const *ptr_B = nullptr;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_C{};
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_C = nullptr;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_D{};
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_D = nullptr;
|
||||
typename EpilogueOutputOp::Params output_op {};
|
||||
int *semaphore = nullptr;
|
||||
SplitKMode split_k_mode {};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(): gemm_k_iterations(0) { }
|
||||
Params() = default;
|
||||
|
||||
///
|
||||
CUTLASS_HOST_DEVICE
|
||||
|
||||
@@ -0,0 +1,494 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 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.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
/*! \file
|
||||
\brief Convolution kernel with an epilogue that computes the absolute maximum value of the output
|
||||
and a pre-activation-function auxiliary output. The auxiliary output is also (optionally)
|
||||
stored to global memory.
|
||||
*/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
|
||||
#include "cutlass/aligned_buffer.h"
|
||||
#include "cutlass/array.h"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/matrix_shape.h"
|
||||
#include "cutlass/semaphore.h"
|
||||
#include "cutlass/tensor_ref.h"
|
||||
#include "cutlass/layout/tensor.h"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/conv2d_problem_size.h"
|
||||
#include "cutlass/conv/conv3d_problem_size.h"
|
||||
#include "cutlass/epilogue/threadblock/output_iterator_parameter.h"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass {
|
||||
namespace conv {
|
||||
namespace kernel {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
typename Mma_, ///! Threadblock-scoped matrix multiply-accumulate
|
||||
typename Epilogue_, ///! Epilogue
|
||||
typename ThreadblockSwizzle_, ///! Threadblock swizzling function
|
||||
conv::Operator ConvOperator, ///! Convolutional operator (Fprop, Dgrad, Wgrad)
|
||||
typename ConvProblemSize_ = Conv2dProblemSize ///! Convolutional operator on 2D or 3D problem
|
||||
>
|
||||
struct ImplicitGemmConvolutionWithAbsMax {
|
||||
|
||||
using Mma = Mma_;
|
||||
using Epilogue = Epilogue_;
|
||||
using EpilogueOutputOp = typename Epilogue::OutputOp;
|
||||
using ThreadblockSwizzle = ThreadblockSwizzle_;
|
||||
static Operator const kConvolutionalOperator = ConvOperator;
|
||||
|
||||
using ElementA = typename Mma::IteratorA::Element;
|
||||
using LayoutA = typename Mma::IteratorA::Layout;
|
||||
using ElementB = typename Mma::IteratorB::Element;
|
||||
using LayoutB = typename Mma::IteratorB::Layout;
|
||||
using ElementC = typename EpilogueOutputOp::ElementOutput;
|
||||
|
||||
/// Set output tensor C layout
|
||||
using LayoutC = LayoutA;
|
||||
|
||||
using ElementAccumulator = typename EpilogueOutputOp::ElementAccumulator;
|
||||
using ElementCompute = typename EpilogueOutputOp::ElementCompute;
|
||||
|
||||
using WarpMmaOperator = typename Mma::Policy::Operator;
|
||||
|
||||
using ArchMmaOperator = typename WarpMmaOperator::ArchMmaOperator;
|
||||
using MathOperator = typename ArchMmaOperator::Operator;
|
||||
|
||||
using OperatorClass = typename WarpMmaOperator::OperatorClass;
|
||||
using ArchTag = typename WarpMmaOperator::ArchTag;
|
||||
|
||||
using ThreadblockShape = typename Mma::Shape;
|
||||
using WarpShape = typename WarpMmaOperator::Shape;
|
||||
using InstructionShape = typename ArchMmaOperator::Shape;
|
||||
|
||||
static int const kStages = Mma::kStages;
|
||||
static IteratorAlgorithm const kIteratorAlgorithm = Mma::IteratorA::kIteratorAlgorithm;
|
||||
static StrideSupport const kStrideSupport = Mma::IteratorA::kStrideSupport;
|
||||
|
||||
/// Warp count (concept: GemmShape)
|
||||
using WarpCount = typename Mma::WarpCount;
|
||||
static int const kThreadCount = 32 * WarpCount::kCount;
|
||||
|
||||
using TensorRefA = typename Mma::IteratorA::TensorRef;
|
||||
using TensorRefB = typename Mma::IteratorB::TensorRef;
|
||||
using TensorRefC = cutlass::TensorRef<ElementC, LayoutC>;
|
||||
using TensorRefAux = cutlass::TensorRef<typename EpilogueOutputOp::ElementAuxOutput, LayoutC>;
|
||||
|
||||
/// Check iterator A and B convolution dimension are the same and
|
||||
// set device::ImplicitGemmConvolution::kConvDim
|
||||
static_assert(Mma::IteratorA::kConvDim == Mma::IteratorB::kConvDim,
|
||||
"Convolution on different different dimensions is not supported");
|
||||
static int const kConvDim = Mma::IteratorA::kConvDim;
|
||||
|
||||
/// Conv dimension and problem size structure (Conv2d or Conv3d)
|
||||
using ConvProblemSize = ConvProblemSize_;
|
||||
|
||||
static conv::GroupMode const kGroupMode = conv::GroupMode::kNone;
|
||||
|
||||
/// Wgrad C stride idx for implicit gemm algorithm
|
||||
// Conv2d row-major matrix C (KxRSC)
|
||||
// Conv3d row-major matrix C (KxTRSC)
|
||||
static int const kWgradCStrideIdx =
|
||||
platform::is_same<LayoutC, cutlass::layout::TensorNHWC>::value ? 2 : 3;
|
||||
|
||||
/// This chooses the appropriate stride element of the C tensor.
|
||||
static int const kTensorCStrideIdx =
|
||||
(kConvolutionalOperator == conv::Operator::kWgrad ? kWgradCStrideIdx : 0);
|
||||
|
||||
//
|
||||
//
|
||||
//
|
||||
using ConvOutputIteratorParameter = epilogue::threadblock::ConvOutputIteratorParameter<
|
||||
LayoutC,
|
||||
typename Epilogue::OutputTileIterator::Layout,
|
||||
TensorRefC,
|
||||
ConvOperator,
|
||||
ConvProblemSize
|
||||
>;
|
||||
|
||||
/// Argument structure
|
||||
struct Arguments {
|
||||
|
||||
//
|
||||
// Data members
|
||||
//
|
||||
|
||||
ConvProblemSize problem_size;
|
||||
TensorRefA ref_A;
|
||||
TensorRefB ref_B;
|
||||
TensorRefC ref_C;
|
||||
TensorRefC ref_D;
|
||||
TensorRefC ref_Aux;
|
||||
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
SplitKMode split_k_mode;
|
||||
|
||||
void * ptr_Vector;
|
||||
|
||||
typename LayoutC::Stride::Index ldr;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
/// Default ctor
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments() { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments(
|
||||
ConvProblemSize const & problem_size
|
||||
):
|
||||
problem_size(problem_size) { }
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments(
|
||||
ConvProblemSize const & problem_size,
|
||||
TensorRefA const & ref_A,
|
||||
TensorRefB const & ref_B,
|
||||
TensorRefC const & ref_C,
|
||||
TensorRefC const & ref_D,
|
||||
TensorRefAux const & ref_Aux,
|
||||
typename EpilogueOutputOp::Params const & output_op,
|
||||
SplitKMode const & split_k_mode = SplitKMode::kSerial,
|
||||
void * ptr_Vector = nullptr,
|
||||
typename LayoutC::Stride::Index ldr = 0
|
||||
):
|
||||
problem_size(problem_size),
|
||||
ref_A(ref_A),
|
||||
ref_B(ref_B),
|
||||
ref_C(ref_C),
|
||||
ref_D(ref_D),
|
||||
ref_Aux(ref_Aux),
|
||||
output_op(output_op),
|
||||
split_k_mode(split_k_mode),
|
||||
ptr_Vector(ptr_Vector),
|
||||
ldr(ldr)
|
||||
{
|
||||
|
||||
}
|
||||
|
||||
};
|
||||
|
||||
/// Parameters structure
|
||||
struct Params {
|
||||
ConvProblemSize problem_size;
|
||||
cutlass::gemm::GemmCoord grid_tiled_shape;
|
||||
gemm::GemmCoord implicit_gemm_problem_size;
|
||||
int swizzle_log_tile;
|
||||
|
||||
int gemm_k_iterations;
|
||||
typename Mma::IteratorA::Params iterator_A;
|
||||
typename Mma::IteratorA::Element const *ptr_A;
|
||||
typename Mma::IteratorB::Params iterator_B;
|
||||
typename Mma::IteratorB::Element const *ptr_B;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_C;
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_C;
|
||||
typename Epilogue::OutputTileIterator::Params iterator_D;
|
||||
typename Epilogue::OutputTileIterator::Element *ptr_D;
|
||||
typename Epilogue::AuxOutputTileIterator::Params iterator_Aux;
|
||||
typename Epilogue::AuxOutputTileIterator::Element *ptr_Aux;
|
||||
typename EpilogueOutputOp::Params output_op;
|
||||
int *semaphore;
|
||||
SplitKMode split_k_mode;
|
||||
|
||||
void * ptr_Vector;
|
||||
typename LayoutC::Stride::Index ldr;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params():
|
||||
swizzle_log_tile(0),
|
||||
gemm_k_iterations(0),
|
||||
ptr_Vector(nullptr),
|
||||
ldr(0)
|
||||
{ }
|
||||
|
||||
///
|
||||
CUTLASS_HOST_DEVICE
|
||||
Params(
|
||||
Arguments const &args,
|
||||
int *semaphore = nullptr
|
||||
):
|
||||
problem_size(args.problem_size),
|
||||
implicit_gemm_problem_size(cutlass::conv::implicit_gemm_problem_size(kConvolutionalOperator, args.problem_size)),
|
||||
iterator_A(Mma::IteratorA::getParams(args.problem_size, args.ref_A.layout())),
|
||||
ptr_A(args.ref_A.data()),
|
||||
iterator_B(args.problem_size, args.ref_B.layout()),
|
||||
ptr_B(args.ref_B.data()),
|
||||
iterator_C(ConvOutputIteratorParameter::layout(args.ref_C)),
|
||||
ptr_C(args.ref_C.data()),
|
||||
iterator_D(ConvOutputIteratorParameter::layout(args.ref_D)),
|
||||
ptr_D(args.ref_D.data()),
|
||||
iterator_Aux(ConvOutputIteratorParameter::layout(args.ref_Aux)),
|
||||
ptr_Aux(args.ref_Aux.data()),
|
||||
output_op(args.output_op),
|
||||
semaphore(semaphore),
|
||||
split_k_mode(args.split_k_mode),
|
||||
ptr_Vector(args.ptr_Vector),
|
||||
ldr(args.ldr)
|
||||
|
||||
{
|
||||
gemm_k_iterations = implicit_gemm_k_iterations(kConvolutionalOperator, ThreadblockShape::kK, args.problem_size);
|
||||
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
grid_tiled_shape = threadblock_swizzle.get_tiled_shape(
|
||||
implicit_gemm_problem_size,
|
||||
{ThreadblockShape::kM, ThreadblockShape::kN, ThreadblockShape::kK},
|
||||
args.problem_size.split_k_slices);
|
||||
|
||||
swizzle_log_tile = threadblock_swizzle.get_log_tile(grid_tiled_shape);
|
||||
}
|
||||
};
|
||||
|
||||
/// Shared memory storage structure
|
||||
union SharedStorage {
|
||||
typename Mma::SharedStorage main_loop;
|
||||
typename Epilogue::SharedStorage epilogue;
|
||||
};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
ImplicitGemmConvolutionWithAbsMax() { }
|
||||
|
||||
/// Executes one ImplicitGEMM
|
||||
CUTLASS_DEVICE
|
||||
void operator()(Params const ¶ms, SharedStorage &shared_storage) {
|
||||
|
||||
// Compute threadblock location
|
||||
ThreadblockSwizzle threadblock_swizzle;
|
||||
|
||||
cutlass::gemm::GemmCoord threadblock_tile_idx =
|
||||
threadblock_swizzle.get_tile_offset(params.swizzle_log_tile);
|
||||
|
||||
// Early exit if CTA is out of range
|
||||
if (params.grid_tiled_shape.m() <= threadblock_tile_idx.m() ||
|
||||
params.grid_tiled_shape.n() <= threadblock_tile_idx.n()) {
|
||||
|
||||
return;
|
||||
}
|
||||
|
||||
// Compute position within threadblock
|
||||
int thread_idx = threadIdx.x;
|
||||
|
||||
// Construct iterators to A and B operands
|
||||
typename Mma::IteratorA iterator_A(
|
||||
params.iterator_A,
|
||||
params.problem_size,
|
||||
params.ptr_A,
|
||||
thread_idx,
|
||||
MatrixCoord(
|
||||
threadblock_tile_idx.m() * Mma::Shape::kM,
|
||||
threadblock_tile_idx.k() * Mma::Shape::kK
|
||||
)
|
||||
);
|
||||
|
||||
typename Mma::IteratorB iterator_B(
|
||||
params.iterator_B,
|
||||
params.problem_size,
|
||||
params.ptr_B,
|
||||
thread_idx,
|
||||
MatrixCoord(
|
||||
threadblock_tile_idx.k() * Mma::Shape::kK,
|
||||
threadblock_tile_idx.n() * Mma::Shape::kN
|
||||
)
|
||||
);
|
||||
|
||||
// Broadcast the warp_id computed by lane 0 to ensure dependent code
|
||||
// is compiled as warp-uniform.
|
||||
int warp_idx = __shfl_sync(0xffffffff, threadIdx.x / 32, 0);
|
||||
int lane_idx = threadIdx.x % 32;
|
||||
|
||||
//
|
||||
// Main loop
|
||||
//
|
||||
|
||||
// Construct thread-scoped matrix multiply
|
||||
Mma mma(shared_storage.main_loop, thread_idx, warp_idx, lane_idx);
|
||||
|
||||
typename Mma::FragmentC accumulators;
|
||||
|
||||
accumulators.clear();
|
||||
|
||||
// Compute threadblock-scoped matrix multiply-add
|
||||
mma(params.gemm_k_iterations, accumulators, iterator_A, iterator_B, accumulators);
|
||||
|
||||
//
|
||||
// Epilogue
|
||||
//
|
||||
|
||||
EpilogueOutputOp output_op(params.output_op);
|
||||
|
||||
// Construct the semaphore.
|
||||
int block_idx = threadblock_tile_idx.m() + threadblock_tile_idx.n() * params.grid_tiled_shape.m();
|
||||
|
||||
Semaphore semaphore(params.semaphore + block_idx, thread_idx);
|
||||
|
||||
// Compute logical position within grid
|
||||
threadblock_tile_idx =
|
||||
threadblock_swizzle.get_tile_offset(params.swizzle_log_tile);
|
||||
|
||||
// If performing a reduction via split-K, fetch the initial synchronization
|
||||
if (params.split_k_mode == SplitKMode::kSerial && params.grid_tiled_shape.k() > 1) {
|
||||
|
||||
// Fetch the synchronization lock initially but do not block.
|
||||
semaphore.fetch();
|
||||
|
||||
// Indicate which position in a serial reduction the output operator is currently updating
|
||||
output_op.set_k_partition(threadblock_tile_idx.k(), params.grid_tiled_shape.k());
|
||||
}
|
||||
|
||||
MatrixCoord threadblock_offset(
|
||||
threadblock_tile_idx.m() * Mma::Shape::kM,
|
||||
threadblock_tile_idx.n() * Mma::Shape::kN
|
||||
);
|
||||
|
||||
// Tile iterator writing to destination tensor
|
||||
typename Epilogue::OutputTileIterator iterator_D(
|
||||
params.iterator_D,
|
||||
params.ptr_D,
|
||||
ConvOutputIteratorParameter::extent(params.problem_size),
|
||||
thread_idx,
|
||||
threadblock_offset
|
||||
);
|
||||
|
||||
// Tile iterator writing to auxiliary tensor.
|
||||
typename Epilogue::AuxOutputTileIterator iterator_Aux(
|
||||
params.iterator_Aux,
|
||||
params.ptr_Aux,
|
||||
ConvOutputIteratorParameter::extent(params.problem_size),
|
||||
thread_idx,
|
||||
threadblock_offset
|
||||
);
|
||||
|
||||
// Tile iterator reading from source accumulator tensor
|
||||
typename Epilogue::OutputTileIterator iterator_C(
|
||||
params.iterator_C,
|
||||
params.ptr_C,
|
||||
ConvOutputIteratorParameter::extent(params.problem_size),
|
||||
thread_idx,
|
||||
threadblock_offset
|
||||
);
|
||||
|
||||
// Define the reduction output pointer and move to the appropriate place
|
||||
typename Epilogue::ElementVector *ptr_Vector =
|
||||
static_cast<typename Epilogue::ElementVector *>(params.ptr_Vector);
|
||||
|
||||
|
||||
// Construct the epilogue
|
||||
Epilogue epilogue(
|
||||
shared_storage.epilogue,
|
||||
thread_idx,
|
||||
warp_idx,
|
||||
lane_idx);
|
||||
|
||||
// Move to appropriate location for this output tile
|
||||
if (ptr_Vector) {
|
||||
ptr_Vector += threadblock_offset.column() + threadblock_tile_idx.m() * params.ldr;
|
||||
}
|
||||
|
||||
// Wait on the semaphore - this latency may have been covered by iterator construction
|
||||
if (params.split_k_mode == SplitKMode::kSerial && params.grid_tiled_shape.k() > 1) {
|
||||
|
||||
// For subsequent threadblocks, the source matrix is held in the 'D' tensor.
|
||||
if (threadblock_tile_idx.k()) {
|
||||
iterator_C = iterator_D;
|
||||
}
|
||||
|
||||
semaphore.wait(threadblock_tile_idx.k());
|
||||
|
||||
}
|
||||
// Each split-k-slice writes to a unique tensor location
|
||||
else if (params.split_k_mode == SplitKMode::kParallel) {
|
||||
iterator_D.add_pointer_offset(threadblock_tile_idx.k() *
|
||||
cutlass::conv::implicit_gemm_tensor_c_size(ConvOperator, params.problem_size));
|
||||
}
|
||||
|
||||
// Execute the epilogue operator to update the destination tensor.
|
||||
epilogue(output_op,
|
||||
// Only the final block uses Vector
|
||||
((params.split_k_mode == SplitKMode::kSerial && params.grid_tiled_shape.k() > 1) &&
|
||||
(params.grid_tiled_shape.k() != threadblock_tile_idx.k() + 1))
|
||||
? nullptr
|
||||
: ptr_Vector,
|
||||
iterator_D,
|
||||
accumulators,
|
||||
iterator_C,
|
||||
iterator_Aux,
|
||||
ConvOutputIteratorParameter::extent(params.problem_size),
|
||||
threadblock_offset);
|
||||
|
||||
//
|
||||
// Release the semaphore
|
||||
//
|
||||
|
||||
if (params.split_k_mode == SplitKMode::kSerial && params.grid_tiled_shape.k() > 1) {
|
||||
|
||||
int lock = 0;
|
||||
if (params.grid_tiled_shape.k() == threadblock_tile_idx.k() + 1) {
|
||||
|
||||
// The final threadblock resets the semaphore for subsequent grids.
|
||||
lock = 0;
|
||||
}
|
||||
else {
|
||||
// Otherwise, the semaphore is incremented
|
||||
lock = threadblock_tile_idx.k() + 1;
|
||||
}
|
||||
|
||||
semaphore.release(lock);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace kernel
|
||||
} // namespace conv
|
||||
} // namespace cutlass
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,391 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2024 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/cluster_sm90.hpp"
|
||||
|
||||
#include "cutlass/conv/convolution.h"
|
||||
#include "cutlass/conv/dispatch_policy.hpp"
|
||||
#include "cutlass/pipeline/sm90_pipeline.hpp"
|
||||
#include "cutlass/gemm/kernel/tile_scheduler.hpp"
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::conv::kernel {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
class CollectiveMainloop_,
|
||||
class CollectiveEpilogue_,
|
||||
class TileSchedulerTag
|
||||
>
|
||||
class ConvUniversal<
|
||||
CollectiveMainloop_,
|
||||
CollectiveEpilogue_,
|
||||
TileSchedulerTag,
|
||||
cute::enable_if_t<cute::is_base_of_v<cutlass::conv::KernelImplicitTmaWarpSpecializedSm90,
|
||||
typename CollectiveMainloop_::DispatchPolicy::Schedule>>>
|
||||
{
|
||||
public:
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
|
||||
// Mainloop derived types
|
||||
using CollectiveMainloop = CollectiveMainloop_;
|
||||
using TileShape = typename CollectiveMainloop::TileShape;
|
||||
using TiledMma = typename CollectiveMainloop::TiledMma;
|
||||
using ArchTag = typename CollectiveMainloop::ArchTag;
|
||||
using ElementA = typename CollectiveMainloop::ElementA;
|
||||
using StrideA = typename CollectiveMainloop::StrideA;
|
||||
using ElementB = typename CollectiveMainloop::ElementB;
|
||||
using StrideB = typename CollectiveMainloop::StrideB;
|
||||
using DispatchPolicy = typename CollectiveMainloop::DispatchPolicy;
|
||||
using ElementAccumulator = typename CollectiveMainloop::ElementAccumulator;
|
||||
using ClusterShape = typename DispatchPolicy::ClusterShape;
|
||||
using MainloopArguments = typename CollectiveMainloop::Arguments;
|
||||
using MainloopParams = typename CollectiveMainloop::Params;
|
||||
static constexpr int NumSpatialDimensions = CollectiveMainloop::NumSpatialDimensions;
|
||||
static_assert(ArchTag::kMinComputeCapability >= 90);
|
||||
|
||||
// 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_assert(cute::is_void_v<TileSchedulerTag>,
|
||||
"TMA warp-specialized kernel does not support specializing the tile scheduler.");
|
||||
using TileScheduler = typename cutlass::gemm::kernel::detail::TileSchedulerSelector<
|
||||
TileSchedulerTag, ArchTag, TileShape, ClusterShape>::Scheduler;
|
||||
using TileSchedulerArguments = typename TileScheduler::Arguments;
|
||||
|
||||
// Kernel level shared memory storage
|
||||
struct SharedStorage {
|
||||
union TensorStorage {
|
||||
using MainloopTensorStorage = typename CollectiveMainloop::TensorStorage;
|
||||
using EpilogueTensorStorage = typename CollectiveEpilogue::TensorStorage;
|
||||
|
||||
MainloopTensorStorage mainloop;
|
||||
EpilogueTensorStorage epilogue;
|
||||
} tensors;
|
||||
|
||||
struct PipelineStorage : cute::aligned_struct<16> {
|
||||
using MainloopPipelineStorage = typename CollectiveMainloop::PipelineStorage;
|
||||
using EpiLoadPipelineStorage = typename CollectiveEpilogue::PipelineStorage;
|
||||
|
||||
alignas(16) MainloopPipelineStorage mainloop;
|
||||
alignas(16) EpiLoadPipelineStorage epi_load;
|
||||
} pipelines;
|
||||
};
|
||||
|
||||
static constexpr int SharedStorageSize = sizeof(SharedStorage);
|
||||
static constexpr uint32_t NumLoadWarpGroups = 1;
|
||||
static constexpr uint32_t NumMmaWarpGroups = 1;
|
||||
static constexpr uint32_t MaxThreadsPerBlock = CUTE_STATIC_V(size(TiledMma{})) + (NumLoadWarpGroups * NumThreadsPerWarpGroup);
|
||||
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
|
||||
|
||||
// Host facing host arguments
|
||||
struct Arguments {
|
||||
MainloopArguments mainloop{};
|
||||
EpilogueArguments epilogue{};
|
||||
KernelHardwareInfo hw_info{};
|
||||
TileSchedulerArguments scheduler{};
|
||||
};
|
||||
|
||||
// Kernel device entry point API
|
||||
struct Params {
|
||||
MainloopParams mainloop;
|
||||
EpilogueParams epilogue;
|
||||
};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
// Map user facing arguments to device facing params
|
||||
static Params
|
||||
to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
(void) workspace;
|
||||
auto mainloop_params = CollectiveMainloop::to_underlying_arguments(args.mainloop, workspace);
|
||||
auto problem_shape_MNKL = append<4>(mainloop_params.problem_shape, Int<1>{});
|
||||
|
||||
return {
|
||||
mainloop_params,
|
||||
CollectiveEpilogue::to_underlying_arguments(problem_shape_MNKL, args.epilogue, workspace)
|
||||
};
|
||||
}
|
||||
|
||||
// Given arguemnts, returns true if the kernel can successfully compute upon them. False otherwise.
|
||||
static bool
|
||||
can_implement(Arguments const& args) {
|
||||
bool implementable = true;
|
||||
implementable &= CollectiveMainloop::can_implement(args.mainloop.problem_shape, args.mainloop);
|
||||
implementable &= CollectiveEpilogue::can_implement(args.mainloop.problem_shape.get_transformed_problem_shape_MNK(), args.epilogue);
|
||||
return implementable;
|
||||
}
|
||||
|
||||
static size_t
|
||||
get_workspace_size(Arguments const& args) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
// Computes the kernel launch grid shape based on runtime parameters
|
||||
static dim3
|
||||
get_grid_shape(Params const& params) {
|
||||
// The CONV mainloop params problem shape will be the cute::Shape<> rank-3 MNK tuple we want for grid planning
|
||||
// Although conv problems do not have an L mode, we add it here to comply with the scheduler API
|
||||
auto linear_problem_shape_MNKL = make_shape(
|
||||
size<0>(params.mainloop.problem_shape), // M mode is linearized.
|
||||
shape<1>(params.mainloop.problem_shape),
|
||||
shape<2>(params.mainloop.problem_shape),
|
||||
Int<1>{});
|
||||
|
||||
return cutlass::gemm::kernel::detail::PersistentTileSchedulerSm90::get_tiled_cta_shape_mnl(
|
||||
linear_problem_shape_MNKL, TileShape{}, ClusterShape{});
|
||||
}
|
||||
|
||||
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;
|
||||
|
||||
// Any Tensor Op MMA Atom in the WGMMA ISA is arch conditional to sm90a.
|
||||
#if ! defined(__CUDA_ARCH_FEAT_SM90_ALL)
|
||||
if constexpr(size<0>(typename TiledMma::AtomShape_MNK{}) == 64) {
|
||||
printf("ERROR : Arch conditional MMA instruction used without targeting sm90a compute capability. Aborting.\n");
|
||||
return;
|
||||
}
|
||||
#endif
|
||||
|
||||
enum class WarpGroupRole {
|
||||
Producer = 0,
|
||||
Consumer = 1,
|
||||
};
|
||||
|
||||
// Kernel level shared memory storage
|
||||
SharedStorage& shared_storage = *reinterpret_cast<SharedStorage*>(smem_buf);
|
||||
|
||||
int thread_idx = int(threadIdx.x);
|
||||
int warp_idx = canonical_warp_idx_sync();
|
||||
int warp_group_thread_idx = thread_idx % NumThreadsPerWarpGroup;
|
||||
auto warp_group_role = WarpGroupRole(canonical_warp_group_idx());
|
||||
int lane_predicate = cute::elect_one_sync();
|
||||
|
||||
// Issue Tma Descriptor Prefetch from a single thread
|
||||
if ((warp_idx == 0) && lane_predicate) {
|
||||
CollectiveMainloop::prefetch_tma_descriptors(params.mainloop);
|
||||
CollectiveEpilogue::prefetch_tma_descriptors(params.epilogue);
|
||||
}
|
||||
|
||||
// Mainloop Load pipeline
|
||||
using MainloopPipeline = typename CollectiveMainloop::MainloopPipeline;
|
||||
typename MainloopPipeline::Params mainloop_pipeline_params;
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Consumer) {
|
||||
mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Consumer;
|
||||
}
|
||||
mainloop_pipeline_params.is_leader = warp_group_thread_idx == 0;
|
||||
mainloop_pipeline_params.num_consumers = NumThreadsPerWarpGroup;
|
||||
mainloop_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
|
||||
MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, mainloop_pipeline_params, ClusterShape{});
|
||||
|
||||
// Epilogue Load pipeline
|
||||
using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline;
|
||||
typename EpiLoadPipeline::Params epi_load_pipeline_params;
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Producer;
|
||||
}
|
||||
if (warp_group_role == WarpGroupRole::Consumer) {
|
||||
epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Consumer;
|
||||
}
|
||||
epi_load_pipeline_params.dst_blockid = cute::block_rank_in_cluster();
|
||||
epi_load_pipeline_params.producer_arv_count = 1; // 1 thread issues TMA load
|
||||
epi_load_pipeline_params.consumer_arv_count = NumThreadsPerWarpGroup;
|
||||
epi_load_pipeline_params.transaction_bytes = CollectiveEpilogue::TmaTransactionBytes;
|
||||
EpiLoadPipeline epi_load_pipeline(shared_storage.pipelines.epi_load, epi_load_pipeline_params);
|
||||
|
||||
// Epilogue Store pipeline
|
||||
using EpiStorePipeline = typename CollectiveEpilogue::StorePipeline;
|
||||
typename EpiStorePipeline::Params epi_store_pipeline_params;
|
||||
epi_store_pipeline_params.always_wait = true;
|
||||
EpiStorePipeline epi_store_pipeline(epi_store_pipeline_params);
|
||||
|
||||
// Initialize starting pipeline states for the collectives
|
||||
// Epilogue store pipe is producer-only (consumer is TMA unit, waits via scoreboarding)
|
||||
typename CollectiveMainloop::PipelineState mainloop_pipe_consumer_state;
|
||||
typename CollectiveEpilogue::LoadPipelineState epi_load_pipe_consumer_state;
|
||||
|
||||
// For the DMA Load (producer) we start with an opposite phase
|
||||
// i.e., we skip all waits since we know that the buffer is indeed empty
|
||||
PipelineState mainloop_pipe_producer_state = cutlass::make_producer_start_state<MainloopPipeline>();
|
||||
PipelineState epi_load_pipe_producer_state = cutlass::make_producer_start_state<EpiLoadPipeline>();
|
||||
PipelineState epi_store_pipe_producer_state = cutlass::make_producer_start_state<EpiStorePipeline>();
|
||||
|
||||
// Separate out problem shape for convenience
|
||||
auto M = get<0>(params.mainloop.problem_shape);
|
||||
auto N = get<1>(params.mainloop.problem_shape);
|
||||
auto K = get<2>(params.mainloop.problem_shape);
|
||||
// output strides are coalesced so we linearize the output shape to match the shape/stride profiles
|
||||
auto linear_problem_shape_MNKL = make_shape(size(M), N, K, Int<1>{});
|
||||
|
||||
// TMA requires special handling of strides to deal with coord codomain mapping
|
||||
// Represent the full tensors -- get these from TMA
|
||||
Tensor mA_mk = params.mainloop.tma_load_a.get_tma_tensor(make_shape(M, size(K)));
|
||||
Tensor mB_nk = params.mainloop.tma_load_b.get_tma_tensor(make_shape(N, K));
|
||||
|
||||
// Get the appropriate blocks for this thread block -- potential for thread block locality
|
||||
auto cta_tile_shape = TileShape{}; // (BLK_M,BLK_N,BLK_K)
|
||||
TiledMma tiled_mma;
|
||||
|
||||
// Make tiled views, defer the slice
|
||||
Tensor gA_mk = local_tile(mA_mk, cta_tile_shape, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M,BLK_K,m,k)
|
||||
Tensor gB_nk = local_tile(mB_nk, cta_tile_shape, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k)
|
||||
|
||||
// Compute m_coord, n_coord, and l_coord with their post-tiled shapes
|
||||
auto m_coord = idx2crd(int(blockIdx.x), shape<2>(gA_mk));
|
||||
auto n_coord = idx2crd(int(blockIdx.y), shape<2>(gB_nk));
|
||||
// The output shape M is linearized so the output coord M here should also be linearized.
|
||||
auto output_tile_coord = make_coord(int(blockIdx.x), n_coord, _, Int<0>{});
|
||||
|
||||
// Slice with m_coord and n_coord
|
||||
Tensor gA = gA_mk(_,_,m_coord,_); // (BLK_M,BLK_K,k)
|
||||
Tensor gB = gB_nk(_,_,n_coord,_); // (BLK_N,BLK_K,k)
|
||||
|
||||
// Get pipeline iterators and increments from tensor shapes
|
||||
auto k_tile_iter = cute::make_coord_iterator(shape<2>(gA));
|
||||
auto k_tile_count = size<2>(gA);
|
||||
|
||||
auto c_tile_count = CollectiveEpilogue::get_load_pipe_increment(cta_tile_shape);
|
||||
auto d_tile_count = CollectiveEpilogue::get_store_pipe_increment(cta_tile_shape);
|
||||
|
||||
// Make sure pipeline init is visible to all producers and consumer CTAs in cluster
|
||||
if constexpr (size(ClusterShape{}) > 1) {
|
||||
cute::cluster_arrive_relaxed();
|
||||
cute::cluster_wait();
|
||||
}
|
||||
else {
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// In a warp specialized kernel, collectives expose data movement and compute operations separately
|
||||
CollectiveMainloop collective_mainloop;
|
||||
CollectiveEpilogue collective_epilogue{params.epilogue, shared_storage.tensors.epilogue};
|
||||
|
||||
if (warp_group_role == WarpGroupRole::Producer) {
|
||||
collective_mainloop.load(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_producer_state,
|
||||
gA, params.mainloop.tma_load_a,
|
||||
gB, params.mainloop.tma_load_b,
|
||||
k_tile_iter, k_tile_count,
|
||||
thread_idx,
|
||||
shared_storage.tensors.mainloop
|
||||
);
|
||||
// Update starting mainloop pipeline state for the pipeline drain
|
||||
mainloop_pipe_producer_state.advance(k_tile_count);
|
||||
// Make sure mainloop consumer has been waited upon before issuing epilogue load
|
||||
collective_mainloop.load_tail(mainloop_pipeline, mainloop_pipe_producer_state);
|
||||
|
||||
if (collective_epilogue.is_producer_load_needed()) {
|
||||
collective_epilogue.load(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_producer_state,
|
||||
linear_problem_shape_MNKL,
|
||||
cta_tile_shape,
|
||||
output_tile_coord,
|
||||
tiled_mma,
|
||||
warp_group_thread_idx,
|
||||
shared_storage.tensors.epilogue
|
||||
);
|
||||
// Update starting load pipeline state for the pipeline drain
|
||||
epi_load_pipe_producer_state.advance(c_tile_count);
|
||||
collective_epilogue.load_tail(epi_load_pipeline, epi_load_pipe_producer_state);
|
||||
}
|
||||
}
|
||||
else if (warp_group_role == WarpGroupRole::Consumer) {
|
||||
Tensor accumulators = partition_fragment_C(tiled_mma, take<0,2>(cta_tile_shape)); // (MMA,MMA_M,MMA_N)
|
||||
|
||||
collective_mainloop.mma(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_consumer_state,
|
||||
accumulators,
|
||||
k_tile_count,
|
||||
thread_idx,
|
||||
shared_storage.tensors.mainloop,
|
||||
params.mainloop
|
||||
);
|
||||
|
||||
// Make sure the math instructions are done and free buffers before entering the epilogue
|
||||
collective_mainloop.mma_tail(
|
||||
mainloop_pipeline,
|
||||
mainloop_pipe_consumer_state,
|
||||
k_tile_count
|
||||
);
|
||||
|
||||
// Epilogue and write to gD
|
||||
collective_epilogue.store(
|
||||
epi_load_pipeline,
|
||||
epi_load_pipe_consumer_state,
|
||||
epi_store_pipeline,
|
||||
epi_store_pipe_producer_state,
|
||||
linear_problem_shape_MNKL,
|
||||
cta_tile_shape,
|
||||
output_tile_coord,
|
||||
accumulators,
|
||||
tiled_mma,
|
||||
warp_group_thread_idx,
|
||||
shared_storage.tensors.epilogue
|
||||
);
|
||||
}
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::conv::kernel
|
||||
@@ -139,7 +139,7 @@ public:
|
||||
// accuracy, where each mainloop iteration first accumulates into a temporary
|
||||
// set of freshly-cleared accumulators, which are subsequently added to the
|
||||
// final accumulator set.
|
||||
static bool const kStagedAccumulation = arch::UseStagedAccumulation<typename Operator::MathOperator>::value;
|
||||
static bool const kStagedAccumulation = arch::detail::UseStagedAccumulation<Operator>::value;
|
||||
};
|
||||
|
||||
private:
|
||||
|
||||
Reference in New Issue
Block a user