v4.1 release
This commit is contained in:
@@ -569,6 +569,47 @@ sm100_make_trivial_fastFP32_tiled_mma() {
|
||||
}
|
||||
}
|
||||
|
||||
template<
|
||||
class CtaShape_MNK
|
||||
>
|
||||
constexpr auto
|
||||
sm100_simt_f32_warp_shape_mnk_selector() {
|
||||
using namespace cute;
|
||||
|
||||
constexpr int CtaShape_M = cute::size<0>(CtaShape_MNK{});
|
||||
constexpr int CtaShape_N = cute::size<1>(CtaShape_MNK{});
|
||||
constexpr int CtaShape_K = cute::size<2>(CtaShape_MNK{});
|
||||
|
||||
// CTA tile shape M and N are supposed to be divisible by 32.
|
||||
static_assert(CtaShape_M % 32 == 0, "CtaShape_M needs to be divisible by 32.");
|
||||
static_assert(CtaShape_N % 32 == 0, "CtaShape_N needs to be divisible by 32.");
|
||||
|
||||
// WarpShape_MNK configuration
|
||||
// We assume WarpShape_K is always 1 in our SM100 SIMT SGEMM implementation.
|
||||
if constexpr (CtaShape_M >= CtaShape_N) {
|
||||
if constexpr (CtaShape_M == 256 && CtaShape_N == 128) {
|
||||
return cute::Shape<_4, _2, _1>{};
|
||||
}
|
||||
else if constexpr ((CtaShape_M == 64 || CtaShape_M == 32) && CtaShape_N == 32) {
|
||||
return cute::Shape<_1, _2, _1>{};
|
||||
}
|
||||
else {
|
||||
return cute::Shape<_2, _2, _1>{};
|
||||
}
|
||||
}
|
||||
else {
|
||||
if constexpr (CtaShape_M == 128 && CtaShape_N == 256) {
|
||||
return cute::Shape<_2, _4, _1>{};
|
||||
}
|
||||
else if constexpr (CtaShape_M == 32 && CtaShape_N == 64) {
|
||||
return cute::Shape<_1, _2, _1>{};
|
||||
}
|
||||
else {
|
||||
return cute::Shape<_1, _4, _1>{};
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
template <
|
||||
class ElementPairA,
|
||||
|
||||
@@ -0,0 +1,216 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/gemm/collective/builders/sm100_common.inl"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::collective {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
template<
|
||||
class LayoutA,
|
||||
int AlignmentA,
|
||||
class LayoutB,
|
||||
int AlignmentB,
|
||||
class CtaShape_MNK,
|
||||
class WarpShape_MNK
|
||||
>
|
||||
constexpr auto
|
||||
sm100_make_simt_f32_tiled_mma() {
|
||||
using namespace cute;
|
||||
|
||||
constexpr int CtaShape_M = cute::size<0>(CtaShape_MNK{});
|
||||
constexpr int CtaShape_N = cute::size<1>(CtaShape_MNK{});
|
||||
constexpr int CtaShape_K = cute::size<2>(CtaShape_MNK{});
|
||||
|
||||
constexpr int WarpShape_M = cute::size<0>(WarpShape_MNK{});
|
||||
constexpr int WarpShape_N = cute::size<1>(WarpShape_MNK{});
|
||||
constexpr int WarpShape_K = cute::size<2>(WarpShape_MNK{});
|
||||
|
||||
// Use Permutation to achieve a [4 x 4] value layout for each thread.
|
||||
// Ideally, we want the tiled mma to be such that loads from shared memory are 128 bit wide.
|
||||
// While as we are using CtaShape_K = 16, when A and B are K-major, we use tranpose + 8 byte padding to avoid smem bank conflict,
|
||||
// so we could only use 64 bit smem load.
|
||||
// When A and B are MN-major, we use 128 bit smem load.
|
||||
using PermutationA = Layout<Shape<_2, Int<WarpShape_M * 8>, _2>, Stride< _1, _4, _2>>;
|
||||
using PermutationB = Layout<Shape<Int<WarpShape_N * 4>, _4>, Stride< _4, _1>>;
|
||||
|
||||
// For 32 threads in 1 warp, we use [8 x 4] thread layouts and each thread will hold [4 x 4] value layouts.
|
||||
// Then totally each warp will hold [32 x 16] value layouts.
|
||||
// So WarpShape_M needs to be equal or smaller than CtaShape_M / 32 and WarpShape_N needs to be equal or smaller than CtaShape_N / 16.
|
||||
static_assert(WarpShape_M <= CtaShape_M / 32, "WarpShape_M is too large, it needs to be equal or smaller than CtaShape_M / 32.");
|
||||
static_assert(WarpShape_N <= CtaShape_N / 16, "WarpShape_N is too large, it needs to be equal or smaller than CtaShape_N / 16.");
|
||||
|
||||
constexpr int WarpStride_M = (WarpShape_M != 1) * NumThreadsPerWarp;
|
||||
constexpr int WarpStride_N = WarpShape_M * NumThreadsPerWarp;
|
||||
|
||||
// We first introduce a [8 x 4] thread layouts in 1 warp.
|
||||
// And inside this [8 x 4] thread layouts, each 4 threads will be arranged as [2 x 2].
|
||||
// Then we could set different WarpShape to finalize how many warps we use in our tiled mma.
|
||||
// For example :
|
||||
// With 128 threads in the tiled mma, we could set the WarpShapeMNK as [2 x 2 x 1], [1 x 4 x 1] and [4 x 1 x 1].
|
||||
// With 64 threads in the tiled mma, we could set the WarpShapeMNK as [1 x 2 x 1] and [2 x 1 x 1].
|
||||
return make_tiled_mma(
|
||||
MMA_Atom<SM100_2x1x1_F32F32F32F32>{},
|
||||
Layout<Shape < Shape <_2, _4, Int<WarpShape_M>>, Shape <_2, _2, Int<WarpShape_N>>, _1>,
|
||||
Stride< Stride<_1, _8, Int<WarpStride_M>>, Stride<_2, _4, Int<WarpStride_N>>, _1>>{},
|
||||
Tile<
|
||||
PermutationA,
|
||||
PermutationB,
|
||||
Underscore>{});
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
template <
|
||||
class GmemLayoutATag,
|
||||
int AlignmentA,
|
||||
class GmemLayoutBTag,
|
||||
int AlignmentB,
|
||||
class CtaShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
int stages,
|
||||
class BuilderScheduleTag>
|
||||
struct CollectiveBuilder<
|
||||
arch::Sm100,
|
||||
arch::OpClassSimt,
|
||||
float,
|
||||
GmemLayoutATag,
|
||||
AlignmentA,
|
||||
float,
|
||||
GmemLayoutBTag,
|
||||
AlignmentB,
|
||||
float,
|
||||
CtaShape_MNK,
|
||||
ClusterShape_MNK,
|
||||
StageCount<stages>,
|
||||
BuilderScheduleTag,
|
||||
cute::enable_if_t<
|
||||
(cute::is_same_v<BuilderScheduleTag, KernelMultistage> ||
|
||||
cute::is_same_v<BuilderScheduleTag, KernelPtrArrayMultistage> ||
|
||||
cute::is_same_v<BuilderScheduleTag, KernelScheduleAuto>) &&
|
||||
((sizeof(float) * AlignmentA) % detail::cp_async_min_alignment_bytes == 0) &&
|
||||
((sizeof(float) * AlignmentB) % detail::cp_async_min_alignment_bytes == 0) >> {
|
||||
static_assert(cute::size<2>(CtaShape_MNK{}) == 16, "SM100 SIMT SGEMM Kernels only support TileShape_K = 16.");
|
||||
|
||||
// This kernel is specialized for F32 data type.
|
||||
using ElementA = float;
|
||||
using ElementB = float;
|
||||
|
||||
using M = decltype(cute::size<0>(CtaShape_MNK{}));
|
||||
using N = decltype(cute::size<1>(CtaShape_MNK{}));
|
||||
using K = decltype(cute::size<2>(CtaShape_MNK{}));
|
||||
|
||||
using WarpShape_MNK = decltype(detail::sm100_simt_f32_warp_shape_mnk_selector<CtaShape_MNK>());
|
||||
|
||||
static constexpr int ThreadCount = cute::size(WarpShape_MNK{}) * NumThreadsPerWarp;
|
||||
|
||||
using TiledMma = decltype(
|
||||
detail::sm100_make_simt_f32_tiled_mma<
|
||||
GmemLayoutATag,
|
||||
AlignmentA,
|
||||
GmemLayoutBTag,
|
||||
AlignmentB,
|
||||
CtaShape_MNK,
|
||||
WarpShape_MNK>());
|
||||
|
||||
// for K major layouts, add a smem alignment offset to avoid bank conflicts
|
||||
static constexpr int SmemAlignmentOffsetA = cutlass::gemm::detail::is_mn_major_A<GmemLayoutATag>() ? 0 : 2;
|
||||
static constexpr int SmemAlignmentOffsetB = cutlass::gemm::detail::is_mn_major_B<GmemLayoutBTag>() ? 0 : 2;
|
||||
static constexpr int CtaShape_M = cute::size<0>(CtaShape_MNK{});
|
||||
static constexpr int CtaShape_N = cute::size<1>(CtaShape_MNK{});
|
||||
|
||||
// Shared memory layout is [M x K] in M-major
|
||||
using SmemLayoutAtomA = cute::Layout<cute::Shape< M, K>,
|
||||
cute::Stride<_1, Int<CtaShape_M + SmemAlignmentOffsetA>>>;
|
||||
// A M-major use 128bit smem load.
|
||||
// A K-major needs to do tranpose and 8 byte padding to make smem bank conflict free, then we can only use 64bit smem load.
|
||||
using SmemCopyAtomA = std::conditional_t<cutlass::gemm::detail::is_mn_major_A<GmemLayoutATag>(),
|
||||
cute::Copy_Atom<cute::AutoVectorizingCopyWithAssumedAlignment<128>, ElementA>,
|
||||
cute::Copy_Atom<cute::AutoVectorizingCopyWithAssumedAlignment<64>, ElementA>>;
|
||||
|
||||
using AlignmentTypeA = cute::uint_byte_t<static_cast<int>(sizeof(ElementA)) * AlignmentA>;
|
||||
using GmemCopyAtomA = cute::Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS_ZFILL<AlignmentTypeA>, ElementA>;
|
||||
using GmemTiledCopyA = decltype(
|
||||
detail::make_simt_gmem_tiled_copy<
|
||||
GmemCopyAtomA, ThreadCount, AlignmentA, TagToStrideA_t<GmemLayoutATag>, M, K>());
|
||||
|
||||
// Shared memory layout is [N x K] in N-major
|
||||
using SmemLayoutAtomB = cute::Layout<cute::Shape< N, K>,
|
||||
cute::Stride<_1, Int<CtaShape_N + SmemAlignmentOffsetB>>>;
|
||||
// B N-major use 128bit smem load.
|
||||
// B K-major needs to do tranpose and 8 byte padding to make smem bank conflict free, then we can only use 64bit smem load.
|
||||
using SmemCopyAtomB = std::conditional_t<cutlass::gemm::detail::is_mn_major_B<GmemLayoutBTag>(),
|
||||
cute::Copy_Atom<cute::AutoVectorizingCopyWithAssumedAlignment<128>, ElementB>,
|
||||
cute::Copy_Atom<cute::AutoVectorizingCopyWithAssumedAlignment<64>, ElementB>>;
|
||||
|
||||
using AlignmentTypeB = cute::uint_byte_t<static_cast<int>(sizeof(ElementB)) * AlignmentB>;
|
||||
using GmemCopyAtomB = cute::Copy_Atom<SM80_CP_ASYNC_CACHEALWAYS_ZFILL<AlignmentTypeB>, ElementB>;
|
||||
using GmemTiledCopyB = decltype(
|
||||
detail::make_simt_gmem_tiled_copy<
|
||||
GmemCopyAtomB, ThreadCount, AlignmentB, TagToStrideB_t<GmemLayoutBTag>, N, K>());
|
||||
|
||||
static constexpr bool IsArrayOfPointersGemm = cute::is_same_v<BuilderScheduleTag, KernelPtrArrayMultistage>;
|
||||
using DispatchPolicy = cute::conditional_t<IsArrayOfPointersGemm,
|
||||
cutlass::gemm::MainloopSm80ArrayCpAsync<stages,
|
||||
ClusterShape_MNK>,
|
||||
cutlass::gemm::MainloopSm80CpAsync<stages,
|
||||
ClusterShape_MNK>
|
||||
>;
|
||||
|
||||
using CollectiveOp = cutlass::gemm::collective::CollectiveMma<
|
||||
DispatchPolicy,
|
||||
CtaShape_MNK,
|
||||
ElementA,
|
||||
TagToStrideA_t<GmemLayoutATag>,
|
||||
ElementB,
|
||||
TagToStrideB_t<GmemLayoutBTag>,
|
||||
TiledMma,
|
||||
GmemTiledCopyA,
|
||||
SmemLayoutAtomA,
|
||||
SmemCopyAtomA,
|
||||
cute::identity,
|
||||
GmemTiledCopyB,
|
||||
SmemLayoutAtomB,
|
||||
SmemCopyAtomB,
|
||||
cute::identity
|
||||
>;
|
||||
};
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -46,6 +46,7 @@
|
||||
#include "cutlass/gemm/collective/builders/sm100_blockscaled_umma_builder.inl"
|
||||
#include "cutlass/gemm/collective/builders/sm100_blockwise_umma_builder.inl"
|
||||
#include "cutlass/gemm/collective/builders/sm100_blockscaled_sparse_umma_builder.inl"
|
||||
#include "cutlass/gemm/collective/builders/sm100_simt_builder.inl"
|
||||
#include "cutlass/gemm/collective/builders/sm120_mma_builder.inl"
|
||||
#include "cutlass/gemm/collective/builders/sm120_blockscaled_mma_builder.inl"
|
||||
#include "cutlass/gemm/collective/builders/sm120_sparse_mma_builder.inl"
|
||||
|
||||
@@ -37,6 +37,7 @@
|
||||
|
||||
#include "cutlass/gemm/collective/sm70_mma_twostage.hpp"
|
||||
#include "cutlass/gemm/collective/sm80_mma_multistage.hpp"
|
||||
#include "cutlass/gemm/collective/sm80_mma_array_multistage.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_multistage_gmma_ss_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_multistage_gmma_rs_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_tma_gmma_ss.hpp"
|
||||
|
||||
@@ -0,0 +1,412 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "cute/algorithm/functional.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cute/algorithm/gemm.hpp"
|
||||
#include "cute/numeric/arithmetic_tuple.hpp"
|
||||
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::collective {
|
||||
using namespace cute;
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
int Stages,
|
||||
class ClusterShape_,
|
||||
class TileShape_,
|
||||
class ElementA_,
|
||||
class StrideA_,
|
||||
class ElementB_,
|
||||
class StrideB_,
|
||||
class TiledMma_,
|
||||
class GmemTiledCopyA_,
|
||||
class SmemLayoutAtomA_,
|
||||
class SmemCopyAtomA_,
|
||||
class TransformA_,
|
||||
class GmemTiledCopyB_,
|
||||
class SmemLayoutAtomB_,
|
||||
class SmemCopyAtomB_,
|
||||
class TransformB_
|
||||
>
|
||||
struct CollectiveMma<
|
||||
MainloopSm80ArrayCpAsync<
|
||||
Stages,
|
||||
ClusterShape_>,
|
||||
TileShape_,
|
||||
ElementA_,
|
||||
StrideA_,
|
||||
ElementB_,
|
||||
StrideB_,
|
||||
TiledMma_,
|
||||
GmemTiledCopyA_,
|
||||
SmemLayoutAtomA_,
|
||||
SmemCopyAtomA_,
|
||||
TransformA_,
|
||||
GmemTiledCopyB_,
|
||||
SmemLayoutAtomB_,
|
||||
SmemCopyAtomB_,
|
||||
TransformB_
|
||||
>
|
||||
{
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
using DispatchPolicy = MainloopSm80ArrayCpAsync<
|
||||
Stages,
|
||||
ClusterShape_>;
|
||||
using TileShape = TileShape_;
|
||||
// Follow the change in TestSmall: TileShape switch to CtaShape
|
||||
// In legacy arch, it should be same
|
||||
using CtaShape_MNK = TileShape;
|
||||
using ElementA = ElementA_;
|
||||
using StrideA = StrideA_;
|
||||
using InternalStrideA = cute::remove_pointer_t<StrideA>;
|
||||
using ElementB = ElementB_;
|
||||
using StrideB = StrideB_;
|
||||
using InternalStrideB = cute::remove_pointer_t<StrideB>;
|
||||
using TiledMma = TiledMma_;
|
||||
using ElementAccumulator = typename TiledMma::ValTypeC; using GmemTiledCopyA = GmemTiledCopyA_;
|
||||
using GmemTiledCopyB = GmemTiledCopyB_;
|
||||
using SmemLayoutAtomA = SmemLayoutAtomA_;
|
||||
using SmemLayoutAtomB = SmemLayoutAtomB_;
|
||||
using SmemCopyAtomA = SmemCopyAtomA_;
|
||||
using SmemCopyAtomB = SmemCopyAtomB_;
|
||||
using TransformA = TransformA_;
|
||||
using TransformB = TransformB_;
|
||||
using ArchTag = typename DispatchPolicy::ArchTag;
|
||||
using ArrayElementA = ElementA;
|
||||
using ArrayElementB = ElementB;
|
||||
static_assert(cute::rank(SmemLayoutAtomA{}) == 2, "SmemLayoutAtom must be rank 2 (M/N, K)");
|
||||
static_assert((size<0>(TileShape{}) % size<0>(SmemLayoutAtomA{})) == 0, "SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert((size<2>(TileShape{}) % size<1>(SmemLayoutAtomA{})) == 0, "SmemLayoutAtom must evenly divide tile shape.");
|
||||
|
||||
static_assert(cute::rank(SmemLayoutAtomB{}) == 2, "SmemLayoutAtom must be rank 2 (M/N, K)");
|
||||
static_assert((size<1>(TileShape{}) % size<0>(SmemLayoutAtomB{})) == 0, "SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert((size<2>(TileShape{}) % size<1>(SmemLayoutAtomB{})) == 0, "SmemLayoutAtom must evenly divide tile shape.");
|
||||
|
||||
using SmemLayoutA = decltype(tile_to_shape(
|
||||
SmemLayoutAtomA{},
|
||||
make_shape(shape<0>(TileShape{}), shape<2>(TileShape{}), Int<DispatchPolicy::Stages>{})));
|
||||
using SmemLayoutB = decltype(tile_to_shape(
|
||||
SmemLayoutAtomB{},
|
||||
make_shape(shape<1>(TileShape{}), shape<2>(TileShape{}), Int<DispatchPolicy::Stages>{})));
|
||||
|
||||
static_assert(DispatchPolicy::Stages >= 2, "CpAsync mainloop must have at least 2 stages in the pipeline.");
|
||||
|
||||
struct SharedStorage
|
||||
{
|
||||
cute::array_aligned<ElementA, cute::cosize_v<SmemLayoutA>> smem_a;
|
||||
cute::array_aligned<ElementB, cute::cosize_v<SmemLayoutB>> smem_b;
|
||||
};
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
ElementA const** ptr_A{nullptr};
|
||||
StrideA dA{};
|
||||
ElementB const** ptr_B{nullptr};
|
||||
StrideB dB{};
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
using Params = Arguments;
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
CollectiveMma() = default;
|
||||
|
||||
template <class ProblemShape>
|
||||
static constexpr Params
|
||||
to_underlying_arguments(ProblemShape const& _, Arguments const& args, void* workspace) {
|
||||
(void) workspace;
|
||||
return args;
|
||||
}
|
||||
|
||||
/// Perform a collective-scoped matrix multiply-accumulate
|
||||
template <
|
||||
class FrgTensorD,
|
||||
class TensorA,
|
||||
class TensorB,
|
||||
class FrgTensorC,
|
||||
class KTileIterator,
|
||||
class ResidueMNK
|
||||
>
|
||||
CUTLASS_DEVICE void
|
||||
operator() (
|
||||
FrgTensorD &accum,
|
||||
TensorA gA, // (BLK_M, BLK_K, K_TILES)
|
||||
TensorB gB, // (BLK_N, BLK_K, K_TILES)
|
||||
FrgTensorC const &src_accum,
|
||||
KTileIterator k_tile_iter, int k_tile_count,
|
||||
ResidueMNK residue_mnk,
|
||||
int thread_idx,
|
||||
char *smem_buf)
|
||||
{
|
||||
using namespace cute;
|
||||
|
||||
static_assert(is_rmem<FrgTensorD>::value, "D tensor must be rmem resident.");
|
||||
static_assert(is_gmem<TensorA>::value, "A tensor must be gmem resident.");
|
||||
static_assert(is_gmem<TensorB>::value, "B tensor must be gmem resident.");
|
||||
static_assert(is_rmem<FrgTensorC>::value, "C tensor must be rmem resident.");
|
||||
static_assert(cute::rank(SmemLayoutA{}) == 3, "Smem layout must be rank 3.");
|
||||
static_assert(cute::rank(SmemLayoutB{}) == 3, "Smem layout must be rank 3.");
|
||||
|
||||
// Construct shared memory tiles
|
||||
SharedStorage& storage = *reinterpret_cast<SharedStorage*>(smem_buf);
|
||||
Tensor sA = make_tensor(make_smem_ptr(storage.smem_a.data()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(storage.smem_b.data()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
|
||||
|
||||
CUTE_STATIC_ASSERT_V(size<0>(gA) == size<0>(sA)); // BLK_M
|
||||
CUTE_STATIC_ASSERT_V(size<1>(gA) == size<1>(sA)); // BLK_K
|
||||
CUTE_STATIC_ASSERT_V(size<0>(gB) == size<0>(sB)); // BLK_N
|
||||
CUTE_STATIC_ASSERT_V(size<1>(gB) == size<1>(sB)); // BLK_K
|
||||
CUTE_STATIC_ASSERT_V(size<1>(sA) == size<1>(sB)); // BLK_K
|
||||
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<2>(sA)); // PIPE
|
||||
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<2>(sB)); // PIPE
|
||||
|
||||
// Shift tensor so residue_k is at origin (Can't read any k_coord < residue_k)
|
||||
// This aligns the tensor with BLK_K for all but the 0th k_tile
|
||||
gA = cute::domain_offset(make_coord(0, get<2>(residue_mnk), 0), gA);
|
||||
gB = cute::domain_offset(make_coord(0, get<2>(residue_mnk), 0), gB);
|
||||
|
||||
// Partition the copying of A and B tiles across the threads
|
||||
GmemTiledCopyA gmem_tiled_copy_A;
|
||||
GmemTiledCopyB gmem_tiled_copy_B;
|
||||
auto gmem_thr_copy_A = gmem_tiled_copy_A.get_slice(thread_idx);
|
||||
auto gmem_thr_copy_B = gmem_tiled_copy_B.get_slice(thread_idx);
|
||||
|
||||
Tensor tAgA = gmem_thr_copy_A.partition_S(gA); // (ACPY,ACPY_M,ACPY_K,k)
|
||||
Tensor tAsA = gmem_thr_copy_A.partition_D(sA); // (ACPY,ACPY_M,ACPY_K,PIPE)
|
||||
Tensor tBgB = gmem_thr_copy_B.partition_S(gB); // (BCPY,BCPY_N,BCPY_K,k)
|
||||
Tensor tBsB = gmem_thr_copy_B.partition_D(sB); // (BCPY,BCPY_N,BCPY_K,PIPE)
|
||||
|
||||
//
|
||||
// PREDICATES
|
||||
//
|
||||
|
||||
// Allocate predicate tensors for m and n
|
||||
Tensor tApA = make_tensor<bool>(make_shape(size<1>(tAsA), size<2>(tAsA)), Stride<_1,_0>{});
|
||||
Tensor tBpB = make_tensor<bool>(make_shape(size<1>(tBsB), size<2>(tBsB)), Stride<_1,_0>{});
|
||||
|
||||
// Construct identity layout for sA and sB
|
||||
Tensor cA = make_identity_tensor(make_shape(size<0>(sA), size<1>(sA))); // (BLK_M,BLK_K) -> (blk_m,blk_k)
|
||||
Tensor cB = make_identity_tensor(make_shape(size<0>(sB), size<1>(sB))); // (BLK_N,BLK_K) -> (blk_n,blk_k)
|
||||
|
||||
// Repeat the partitioning with identity layouts
|
||||
Tensor tAcA = gmem_thr_copy_A.partition_S(cA); // (ACPY,ACPY_M,ACPY_K) -> (blk_m,blk_k)
|
||||
Tensor tBcB = gmem_thr_copy_B.partition_S(cB); // (BCPY,BCPY_N,BCPY_K) -> (blk_n,blk_k)
|
||||
|
||||
// Set predicates for m bounds
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int m = 0; m < size<0>(tApA); ++m) {
|
||||
tApA(m,0) = get<0>(tAcA(0,m,0)) < get<0>(residue_mnk); // blk_m coord < residue_m
|
||||
}
|
||||
// Set predicates for n bounds
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int n = 0; n < size<0>(tBpB); ++n) {
|
||||
tBpB(n,0) = get<0>(tBcB(0,n,0)) < get<1>(residue_mnk); // blk_n coord < residue_n
|
||||
}
|
||||
|
||||
//
|
||||
// PREFETCH
|
||||
//
|
||||
|
||||
// Clear the smem tiles to account for predicated off loads
|
||||
clear(tAsA);
|
||||
clear(tBsB);
|
||||
|
||||
// Start async loads for 0th k-tile, where we take care of the k residue
|
||||
{
|
||||
constexpr int k_pipe = 0;
|
||||
|
||||
Tensor tAgAk = tAgA(_,_,_,*k_tile_iter);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k = 0; k < size<2>(tAsA); ++k) {
|
||||
if (get<1>(tAcA(0,0,k)) >= -get<2>(residue_mnk)) { // blk_k coord < residue_k (gA shifted)
|
||||
copy_if(gmem_tiled_copy_A, tApA(_,k), tAgAk(_,_,k), tAsA(_,_,k,k_pipe));
|
||||
}
|
||||
}
|
||||
Tensor tBgBk = tBgB(_,_,_,*k_tile_iter);
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k = 0; k < size<2>(tBsB); ++k) {
|
||||
if (get<1>(tBcB(0,0,k)) >= -get<2>(residue_mnk)) { // blk_k coord < residue_k (gB shifted)
|
||||
copy_if(gmem_tiled_copy_B, tBpB(_,k), tBgBk(_,_,k), tBsB(_,_,k,k_pipe));
|
||||
}
|
||||
}
|
||||
cp_async_fence();
|
||||
++k_tile_iter;
|
||||
--k_tile_count;
|
||||
}
|
||||
|
||||
// Start async loads for 1st k-tile onwards, no k-residue handling needed
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_pipe = 1; k_pipe < DispatchPolicy::Stages-1; ++k_pipe) {
|
||||
if (k_tile_count <= 0) {
|
||||
clear(tApA);
|
||||
clear(tBpB);
|
||||
}
|
||||
copy_if(gmem_tiled_copy_A, tApA, tAgA(_,_,_,*k_tile_iter), tAsA(_,_,_,k_pipe)); // CpAsync
|
||||
copy_if(gmem_tiled_copy_B, tBpB, tBgB(_,_,_,*k_tile_iter), tBsB(_,_,_,k_pipe)); // CpAsync
|
||||
cp_async_fence();
|
||||
++k_tile_iter;
|
||||
--k_tile_count;
|
||||
}
|
||||
|
||||
//
|
||||
// MMA Atom partitioning
|
||||
//
|
||||
|
||||
// Tile MMA compute thread partitions and allocate accumulators
|
||||
TiledMma tiled_mma;
|
||||
auto thr_mma = tiled_mma.get_thread_slice(thread_idx);
|
||||
Tensor tCrA = thr_mma.partition_fragment_A(sA(_,_,0)); // (MMA,MMA_M,MMA_K)
|
||||
Tensor tCrB = thr_mma.partition_fragment_B(sB(_,_,0)); // (MMA,MMA_N,MMA_K)
|
||||
|
||||
CUTE_STATIC_ASSERT_V(size<1>(tCrA) == size<1>(accum)); // MMA_M
|
||||
CUTE_STATIC_ASSERT_V(size<1>(tCrA) == size<1>(src_accum)); // MMA_M
|
||||
CUTE_STATIC_ASSERT_V(size<1>(tCrB) == size<2>(accum)); // MMA_N
|
||||
CUTE_STATIC_ASSERT_V(size<1>(tCrB) == size<2>(src_accum)); // MMA_N
|
||||
CUTE_STATIC_ASSERT_V(size<2>(tCrA) == size<2>(tCrB)); // MMA_K
|
||||
|
||||
//
|
||||
// Copy Atom retiling
|
||||
//
|
||||
|
||||
auto smem_tiled_copy_A = make_tiled_copy_A(SmemCopyAtomA{}, tiled_mma);
|
||||
auto smem_thr_copy_A = smem_tiled_copy_A.get_thread_slice(thread_idx);
|
||||
Tensor tCsA = smem_thr_copy_A.partition_S(sA); // (CPY,CPY_M,CPY_K,PIPE)
|
||||
Tensor tCrA_copy_view = smem_thr_copy_A.retile_D(tCrA); // (CPY,CPY_M,CPY_K)
|
||||
CUTE_STATIC_ASSERT_V(size<1>(tCsA) == size<1>(tCrA_copy_view)); // CPY_M
|
||||
CUTE_STATIC_ASSERT_V(size<2>(tCsA) == size<2>(tCrA_copy_view)); // CPY_K
|
||||
|
||||
auto smem_tiled_copy_B = make_tiled_copy_B(SmemCopyAtomB{}, tiled_mma);
|
||||
auto smem_thr_copy_B = smem_tiled_copy_B.get_thread_slice(thread_idx);
|
||||
Tensor tCsB = smem_thr_copy_B.partition_S(sB); // (CPY,CPY_N,CPY_K,PIPE)
|
||||
Tensor tCrB_copy_view = smem_thr_copy_B.retile_D(tCrB); // (CPY,CPY_N,CPY_K)
|
||||
CUTE_STATIC_ASSERT_V(size<1>(tCsB) == size<1>(tCrB_copy_view)); // CPY_N
|
||||
CUTE_STATIC_ASSERT_V(size<2>(tCsB) == size<2>(tCrB_copy_view)); // CPY_K
|
||||
|
||||
//
|
||||
// PIPELINED MAIN LOOP
|
||||
//
|
||||
|
||||
// Current pipe index in smem to read from
|
||||
int smem_pipe_read = 0;
|
||||
// Current pipe index in smem to write to
|
||||
int smem_pipe_write = DispatchPolicy::Stages-1;
|
||||
|
||||
Tensor tCsA_p = tCsA(_,_,_,smem_pipe_read);
|
||||
Tensor tCsB_p = tCsB(_,_,_,smem_pipe_read);
|
||||
|
||||
// Size of the register pipeline
|
||||
auto K_BLOCK_MAX = size<2>(tCrA);
|
||||
|
||||
// PREFETCH register pipeline
|
||||
if (K_BLOCK_MAX > 1) {
|
||||
// Wait until our first prefetched tile is loaded in
|
||||
cp_async_wait<DispatchPolicy::Stages-2>();
|
||||
__syncthreads();
|
||||
|
||||
// Prefetch the first rmem from the first k-tile
|
||||
copy(smem_tiled_copy_A, tCsA_p(_,_,Int<0>{}), tCrA_copy_view(_,_,Int<0>{}));
|
||||
copy(smem_tiled_copy_B, tCsB_p(_,_,Int<0>{}), tCrB_copy_view(_,_,Int<0>{}));
|
||||
}
|
||||
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
for ( ; k_tile_count > -(DispatchPolicy::Stages-1); --k_tile_count)
|
||||
{
|
||||
// Pipeline the outer products with a static for loop.
|
||||
//
|
||||
// Note, the for_each() function is required here to ensure `k_block` is of type Int<N>.
|
||||
for_each(make_int_sequence<K_BLOCK_MAX>{}, [&] (auto k_block)
|
||||
{
|
||||
if (k_block == K_BLOCK_MAX - 1)
|
||||
{
|
||||
// Slice the smem_pipe_read smem
|
||||
tCsA_p = tCsA(_,_,_,smem_pipe_read);
|
||||
tCsB_p = tCsB(_,_,_,smem_pipe_read);
|
||||
|
||||
// Commit the smem for smem_pipe_read
|
||||
cp_async_wait<DispatchPolicy::Stages-2>();
|
||||
__syncthreads();
|
||||
}
|
||||
|
||||
// Load A, B shmem->regs for k_block+1
|
||||
auto k_block_next = (k_block + Int<1>{}) % K_BLOCK_MAX; // static
|
||||
copy(smem_tiled_copy_A, tCsA_p(_,_,k_block_next), tCrA_copy_view(_,_,k_block_next));
|
||||
copy(smem_tiled_copy_B, tCsB_p(_,_,k_block_next), tCrB_copy_view(_,_,k_block_next));
|
||||
// Copy gmem to smem before computing gemm on each k-pipe
|
||||
if (k_block == 0)
|
||||
{
|
||||
// Set all predicates to false if we are going to overshoot bounds
|
||||
if (k_tile_count <= 0) {
|
||||
clear(tApA);
|
||||
clear(tBpB);
|
||||
}
|
||||
copy_if(gmem_tiled_copy_A, tApA, tAgA(_,_,_,*k_tile_iter), tAsA(_,_,_,smem_pipe_write));
|
||||
copy_if(gmem_tiled_copy_B, tBpB, tBgB(_,_,_,*k_tile_iter), tBsB(_,_,_,smem_pipe_write));
|
||||
cp_async_fence();
|
||||
++k_tile_iter;
|
||||
|
||||
// Advance the pipe -- Doing it here accounts for K_BLOCK_MAX = 1 (no rmem pipe)
|
||||
smem_pipe_write = smem_pipe_read;
|
||||
++smem_pipe_read;
|
||||
smem_pipe_read = (smem_pipe_read == DispatchPolicy::Stages) ? 0 : smem_pipe_read;
|
||||
}
|
||||
|
||||
// Transform before compute
|
||||
cute::transform(tCrA(_,_,k_block), TransformA{});
|
||||
cute::transform(tCrB(_,_,k_block), TransformB{});
|
||||
// Thread-level register gemm for k_block
|
||||
cute::gemm(tiled_mma, accum, tCrA(_,_,k_block), tCrB(_,_,k_block), src_accum);
|
||||
});
|
||||
|
||||
}
|
||||
|
||||
cp_async_wait<0>();
|
||||
__syncthreads();
|
||||
}
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
+1
-6
@@ -89,15 +89,10 @@ struct CollectiveMma<
|
||||
TransformB_>
|
||||
{
|
||||
public:
|
||||
enum class ConversionMode {
|
||||
DirectConvert,
|
||||
ConvertAndScale,
|
||||
ConvertAndScaleWithZero
|
||||
};
|
||||
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
using ConversionMode = cutlass::detail::ConversionMode;
|
||||
using DispatchPolicy = MainloopSm90ArrayTmaGmmaWarpSpecializedMixedInput<Stages, ClusterShape, KernelSchedule_>;
|
||||
using TileShape = TileShape_;
|
||||
using KernelSchedule = KernelSchedule_;
|
||||
|
||||
+1
-5
@@ -96,15 +96,11 @@ struct CollectiveMma<
|
||||
TransformB_>
|
||||
{
|
||||
public:
|
||||
enum class ConversionMode {
|
||||
DirectConvert,
|
||||
ConvertAndScale,
|
||||
ConvertAndScaleWithZero
|
||||
};
|
||||
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
using ConversionMode = cutlass::detail::ConversionMode;
|
||||
using DispatchPolicy = MainloopSm90TmaGmmaRmemAWarpSpecializedMixedInput<Stages, ClusterShape, KernelSchedule_>;
|
||||
using TileShape = TileShape_;
|
||||
using KernelSchedule = KernelSchedule_;
|
||||
|
||||
@@ -109,6 +109,7 @@ static constexpr bool HasAuxiliaryLoad_v = HasAuxiliaryLoad<T>::value;
|
||||
// Kernel schedule policies (the base class tags, one for each kernel layer file)
|
||||
//
|
||||
struct KernelMultistage { };
|
||||
struct KernelPtrArrayMultistage { };
|
||||
struct KernelCpAsyncWarpSpecialized { };
|
||||
struct KernelCpAsyncWarpSpecializedPingpong { };
|
||||
struct KernelCpAsyncWarpSpecializedCooperative { };
|
||||
@@ -198,6 +199,17 @@ struct MainloopSm80CpAsync {
|
||||
using ClusterShape = ClusterShape_;
|
||||
};
|
||||
|
||||
// n-buffer in smem (cp.async), pipelined with registers, with predicated gmem loads for SM100 Simt Ptr-Array
|
||||
template<int Stages_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>
|
||||
>
|
||||
struct MainloopSm80ArrayCpAsync {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ArchTag = cute::conditional_t<(size(ClusterShape_{}) > 1), arch::Sm90, arch::Sm80>;
|
||||
using Schedule = KernelPtrArrayMultistage;
|
||||
using ClusterShape = ClusterShape_;
|
||||
};
|
||||
|
||||
// n-buffer in smem (cp.async), pipelined with Hopper GMMA, with predicated gmem loads, warp specialized dynamic schedule
|
||||
template<
|
||||
int Stages_,
|
||||
@@ -479,6 +491,16 @@ struct KernelTmaWarpSpecializedInputTransformSm100 final {
|
||||
static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_;
|
||||
};
|
||||
|
||||
// InputTransform GEMM
|
||||
template<
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_
|
||||
>
|
||||
struct KernelTmaWarpSpecializedMixedInputTransformSm100 final {
|
||||
static constexpr int SchedulerPipelineStageCount = SchedulerPipelineStageCount_;
|
||||
static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_;
|
||||
};
|
||||
|
||||
// Ptr-Array Dense GEMM: SM100 tensor op policy that applies to both 1SM and 2SM MMA atoms
|
||||
template<
|
||||
int SchedulerPipelineStageCount_,
|
||||
|
||||
@@ -54,6 +54,7 @@ struct IsCutlass3ArrayKernel<ProblemShape, cute::void_t<typename ProblemShape::U
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
#include "cutlass/gemm/kernel/sm70_gemm.hpp"
|
||||
#include "cutlass/gemm/kernel/sm70_gemm_array.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_gemm_tma.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_gemm_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_gemm_warpspecialized_pingpong.hpp"
|
||||
|
||||
@@ -1008,39 +1008,6 @@ public:
|
||||
// Advance the mm2accum pipe
|
||||
mma2accum_pipeline_consumer_state = mma2accum_pipeline_consumer_state_next;
|
||||
}
|
||||
else if constexpr (InputTransformType == cutlass::gemm::detail::KernelInputTransformType::MixedInput) {
|
||||
|
||||
mma2accum_pipeline.consumer_wait(mma2accum_pipeline_consumer_state);
|
||||
|
||||
// Accumulators
|
||||
Tensor accumulators = bulk_tmem(_,_,_,mma2accum_pipeline_consumer_state.index()); // ((MMA_TILE_M,MMA_TILE_N),MMA_M,MMA_N)
|
||||
|
||||
mma2accum_pipeline_consumer_state = scheduler.template fixup<IsComplex>(
|
||||
TiledMma{},
|
||||
work_tile_info,
|
||||
accumulators,
|
||||
mma2accum_pipeline,
|
||||
mma2accum_pipeline_consumer_state,
|
||||
typename CollectiveEpilogue::CopyOpT2R{}
|
||||
);
|
||||
|
||||
//
|
||||
// Epilogue and write to gD
|
||||
//
|
||||
if (scheduler.compute_epilogue(work_tile_info)) {
|
||||
auto [mma2accum_pipeline_state_next] = collective_epilogue(
|
||||
mma2accum_pipeline,
|
||||
mma2accum_pipeline_consumer_state,
|
||||
problem_shape_MNKL,
|
||||
CtaShape_MNK{},
|
||||
cta_coord_mnkl,
|
||||
accumulators,
|
||||
shared_storage.tensors.epilogue
|
||||
);
|
||||
// Advance the mma2accum pipe
|
||||
mma2accum_pipeline_consumer_state = mma2accum_pipeline_state_next;
|
||||
}
|
||||
}
|
||||
// Complex kernels use a collective epilogue
|
||||
else {
|
||||
mma2accum_pipeline.consumer_wait(mma2accum_pipeline_consumer_state);
|
||||
|
||||
@@ -0,0 +1,279 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2023 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
#pragma once
|
||||
|
||||
#include "cutlass/cutlass.h"
|
||||
#include "cutlass/kernel_hardware_info.hpp"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
namespace cutlass::gemm::kernel {
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
class ProblemShape_,
|
||||
class CollectiveMainloop_,
|
||||
class CollectiveEpilogue_,
|
||||
class TileScheduler_
|
||||
>
|
||||
class GemmUniversal<
|
||||
ProblemShape_,
|
||||
CollectiveMainloop_,
|
||||
CollectiveEpilogue_,
|
||||
TileScheduler_,
|
||||
cute::enable_if_t<cute::is_base_of_v<KernelPtrArrayMultistage, typename CollectiveMainloop_::DispatchPolicy::Schedule>>>
|
||||
{
|
||||
public:
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
using ProblemShape = ProblemShape_;
|
||||
static_assert(rank(typename ProblemShape::UnderlyingProblemShape{}) == 4,
|
||||
"ProblemShape{} should be <M,N,K> or <M,N,K,L>");
|
||||
|
||||
// Mainloop derived types
|
||||
using CollectiveMainloop = CollectiveMainloop_;
|
||||
using TileShape = typename CollectiveMainloop::TileShape;
|
||||
using TiledMma = typename CollectiveMainloop::TiledMma;
|
||||
using ArchTag = typename CollectiveMainloop::ArchTag;
|
||||
using ElementA = typename CollectiveMainloop::ElementA;
|
||||
using StrideA = typename CollectiveMainloop::StrideA;
|
||||
using InternalStrideA = typename CollectiveMainloop::InternalStrideA;
|
||||
using ElementB = typename CollectiveMainloop::ElementB;
|
||||
using StrideB = typename CollectiveMainloop::StrideB;
|
||||
using InternalStrideB = typename CollectiveMainloop::InternalStrideB;
|
||||
using DispatchPolicy = typename CollectiveMainloop::DispatchPolicy;
|
||||
using ElementAccumulator = typename CollectiveMainloop::ElementAccumulator;
|
||||
using MainloopArguments = typename CollectiveMainloop::Arguments;
|
||||
using MainloopParams = typename CollectiveMainloop::Params;
|
||||
|
||||
using TileSchedulerTag = TileScheduler_;
|
||||
using TileScheduler = typename detail::TileSchedulerSelector<
|
||||
TileScheduler_, ArchTag, TileShape,
|
||||
cute::Shape<cute::Int<1>, cute::Int<1>, cute::Int<1>>>::Scheduler;
|
||||
using TileSchedulerArguments = typename TileScheduler::Arguments;
|
||||
static constexpr bool IsGdcEnabled = false;
|
||||
|
||||
static constexpr bool is_valid_tile_scheduler =
|
||||
cute::is_void_v<TileScheduler_> or cute::is_same_v<TileScheduler_, PersistentScheduler>;
|
||||
static_assert(is_valid_tile_scheduler, "SM70 kernel does not support specializing the tile scheduler.");
|
||||
|
||||
// Epilogue derived types
|
||||
using CollectiveEpilogue = CollectiveEpilogue_;
|
||||
using ElementC = typename CollectiveEpilogue::ElementC;
|
||||
using StrideC = typename CollectiveEpilogue::StrideC;
|
||||
using InternalStrideC = typename CollectiveEpilogue::InternalStrideC;
|
||||
using ElementD = typename CollectiveEpilogue::ElementD;
|
||||
using StrideD = typename CollectiveEpilogue::StrideD;
|
||||
using InternalStrideD = typename CollectiveEpilogue::InternalStrideD;
|
||||
using EpilogueArguments = typename CollectiveEpilogue::Arguments;
|
||||
using EpilogueParams = typename CollectiveEpilogue::Params;
|
||||
static_assert(cute::is_same_v<ElementAccumulator, typename CollectiveEpilogue::ElementAccumulator>,
|
||||
"Mainloop and epilogue do not agree on accumulator value type.");
|
||||
|
||||
// MSVC requires the cast to fix a warning-as-error.
|
||||
static constexpr int SharedStorageSize = static_cast<int>(cute::max(
|
||||
sizeof(typename CollectiveMainloop::SharedStorage),
|
||||
sizeof(typename CollectiveEpilogue::SharedStorage)));
|
||||
|
||||
static constexpr uint32_t MaxThreadsPerBlock = CUTE_STATIC_V(cute::size(TiledMma{}));
|
||||
static constexpr uint32_t MinBlocksPerMultiprocessor = 1;
|
||||
|
||||
// Device side arguments
|
||||
struct Arguments {
|
||||
GemmUniversalMode mode{};
|
||||
ProblemShape problem_shape{};
|
||||
MainloopArguments mainloop{};
|
||||
EpilogueArguments epilogue{};
|
||||
KernelHardwareInfo hw_info{};
|
||||
TileSchedulerArguments scheduler{};
|
||||
};
|
||||
|
||||
// Kernel entry point API
|
||||
struct Params {
|
||||
GemmUniversalMode mode{};
|
||||
typename ProblemShape::UnderlyingProblemShape problem_shape{};
|
||||
MainloopParams mainloop{};
|
||||
EpilogueParams epilogue{};
|
||||
};
|
||||
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
|
||||
// Convert to underlying arguments. In this case, a simple copy for the aliased type.
|
||||
static
|
||||
Params
|
||||
to_underlying_arguments(Arguments const& args, void* workspace) {
|
||||
(void) workspace;
|
||||
typename ProblemShape::UnderlyingProblemShape problem_shape = args.problem_shape.get_host_problem_shape();
|
||||
|
||||
KernelHardwareInfo hw_info{args.hw_info.device_id, args.hw_info.sm_count};
|
||||
auto problem_shape_MNKL = append<4>(args.problem_shape, Int<1>{});
|
||||
|
||||
return {
|
||||
args.mode,
|
||||
problem_shape,
|
||||
CollectiveMainloop::to_underlying_arguments(problem_shape, args.mainloop, workspace),
|
||||
CollectiveEpilogue::to_underlying_arguments(problem_shape, args.epilogue, workspace)
|
||||
};
|
||||
}
|
||||
|
||||
static bool
|
||||
can_implement(Arguments const& args) {
|
||||
|
||||
bool implementable = (args.mode == GemmUniversalMode::kArray && rank(typename ProblemShape::UnderlyingProblemShape{}) == 4);
|
||||
if (!implementable) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Arguments or Problem Shape don't meet the requirements.\n");
|
||||
return implementable;
|
||||
}
|
||||
typename ProblemShape::UnderlyingProblemShape problem_shape = args.problem_shape.get_host_problem_shape();
|
||||
implementable &= TileScheduler::can_implement(args.scheduler);
|
||||
return implementable;
|
||||
}
|
||||
|
||||
static size_t
|
||||
get_workspace_size(Arguments const& args) {
|
||||
size_t workspace_size = 0;
|
||||
return workspace_size;
|
||||
}
|
||||
|
||||
static
|
||||
cutlass::Status
|
||||
initialize_workspace(Arguments const& args, void* workspace = nullptr, cudaStream_t stream = nullptr,
|
||||
CudaHostAdapter* cuda_adapter = nullptr) {
|
||||
cutlass::Status status = Status::kSuccess;
|
||||
|
||||
return status;
|
||||
}
|
||||
|
||||
static dim3
|
||||
get_grid_shape(Params const& params) {
|
||||
int batch_count = cute::size<3>(params.problem_shape);
|
||||
return dim3(
|
||||
cute::size(cute::ceil_div(cute::shape<0>(params.problem_shape), cute::shape<0>(TileShape{}))),
|
||||
cute::size(cute::ceil_div(cute::shape<1>(params.problem_shape), cute::shape<1>(TileShape{}))),
|
||||
batch_count
|
||||
);
|
||||
}
|
||||
|
||||
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;
|
||||
|
||||
// Preconditions
|
||||
CUTE_STATIC_ASSERT(is_static<TileShape>::value);
|
||||
|
||||
// Separate out problem shape for convenience
|
||||
// Optionally append 1s until problem shape is rank-4 in case its is only rank-3 (MNK)
|
||||
auto problem_shape_MNKL = append<4>(params.problem_shape, Int<1>{});
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
|
||||
// Preconditions
|
||||
static_assert(cute::rank(StrideA{}) == 3, "StrideA must be rank-3: [M, K, L]. If batch mode is not needed, set L stride to Int<0>.");
|
||||
static_assert(cute::rank(StrideB{}) == 3, "StrideB must be rank-3: [N, K, L]. If batch mode is not needed, set L stride to Int<0>.");
|
||||
static_assert(cute::rank(StrideC{}) == 3, "StrideC must be rank-3: [M, N, L]. If batch mode is not needed, set L stride to Int<0>.");
|
||||
static_assert(cute::rank(StrideD{}) == 3, "StrideD must be rank-3: [M, N, L]. If batch mode is not needed, set L stride to Int<0>.");
|
||||
|
||||
// Get the appropriate blocks for this thread block -- potential for thread block locality
|
||||
int thread_idx = int(threadIdx.x);
|
||||
auto blk_shape = TileShape{}; // (BLK_M,BLK_N,BLK_K)
|
||||
auto [m_coord, n_coord, l_coord] = static_cast<uint3>(blockIdx);
|
||||
auto blk_coord_mnkl = make_coord(int(m_coord), int(n_coord), _, int(l_coord)); // (m,n,k,l)
|
||||
|
||||
// Represent the full tensors
|
||||
Tensor mA_mkl = make_tensor(make_gmem_ptr(params.mainloop.ptr_A[l_coord]), make_shape(M,K,1), params.mainloop.dA); //(m,k,l)
|
||||
Tensor mB_nkl = make_tensor(make_gmem_ptr(params.mainloop.ptr_B[l_coord]), make_shape(N,K,1), params.mainloop.dB); //(n,k,l)
|
||||
|
||||
// Get batch slice
|
||||
Tensor mA_mk = mA_mkl(_,_,0); // (m,k)
|
||||
Tensor mB_nk = mB_nkl(_,_,0); // (n,k)
|
||||
|
||||
// Slice to get the tiles this thread block is responsible for
|
||||
Tensor gA = local_tile(mA_mk, blk_shape, take<0,3>(blk_coord_mnkl), Step<_1, X,_1>{}); // (BLK_M,BLK_K,k)
|
||||
Tensor gB = local_tile(mB_nk, blk_shape, take<0,3>(blk_coord_mnkl), Step< X,_1,_1>{}); // (BLK_N,BLK_K,k)
|
||||
|
||||
// Compute tile residues for predication
|
||||
auto m_max_coord = M - size<0>(gA) * get<0>(blk_coord_mnkl); // M - BLK_M * m_coord
|
||||
auto n_max_coord = N - size<0>(gB) * get<1>(blk_coord_mnkl); // N - BLK_N * n_coord
|
||||
auto k_residue = K - size<1>(gA) * size<2>(gA); // K - BLK_K * k_coord_max
|
||||
auto residue_mnk = make_tuple(m_max_coord, n_max_coord, k_residue);
|
||||
|
||||
// Allocate the tiled_mma and the accumulators for the (M,N) blk_shape
|
||||
TiledMma tiled_mma;
|
||||
Tensor accumulators = partition_fragment_C(tiled_mma, take<0,2>(blk_shape)); // (MMA,MMA_M,MMA_N)
|
||||
clear(accumulators);
|
||||
|
||||
auto k_tile_iter = cute::make_coord_iterator(shape<2>(gA));
|
||||
int k_tile_count = size<2>(gA);
|
||||
|
||||
|
||||
// Perform the collective scoped MMA
|
||||
CollectiveMainloop collective_mma;
|
||||
collective_mma(
|
||||
accumulators,
|
||||
gA,
|
||||
gB,
|
||||
accumulators,
|
||||
k_tile_iter, k_tile_count,
|
||||
residue_mnk,
|
||||
thread_idx,
|
||||
smem_buf
|
||||
);
|
||||
|
||||
// Epilogue and write to gD
|
||||
CollectiveEpilogue epilogue{params.epilogue};
|
||||
epilogue(
|
||||
problem_shape_MNKL,
|
||||
blk_shape,
|
||||
blk_coord_mnkl,
|
||||
accumulators,
|
||||
tiled_mma,
|
||||
residue_mnk,
|
||||
thread_idx,
|
||||
smem_buf
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::kernel
|
||||
Reference in New Issue
Block a user