CUTLASS 3.8 Release (#2059)
* CUTLASS 3.8 Release * update * Update README.md * Revert "Update README.md" This reverts commit b353e36fe83e0815f99b44e46c0c95494c44726b. * update * update --------- Co-authored-by: Haicheng Wu <57973641+hwu36@users.noreply.github.com> Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
This commit is contained in:
co-authored by
Haicheng Wu
Haicheng Wu
parent
9eb01fa0b0
commit
389e493055
@@ -0,0 +1,782 @@
|
||||
/***************************************************************************************************
|
||||
* 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"
|
||||
#include "cutlass/gemm/collective/builders/sm100_pipeline_carveout.inl"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::collective {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
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,
|
||||
class TileShapeSFA,
|
||||
class TileShapeSFB,
|
||||
int stages
|
||||
>
|
||||
constexpr int
|
||||
sm100_compute_stage_count_or_override_blockscaled(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,
|
||||
class TileShapeSFA,
|
||||
class TileShapeSFB,
|
||||
int carveout_bytes
|
||||
>
|
||||
constexpr int
|
||||
sm100_compute_stage_count_or_override_blockscaled(StageCountAutoCarveout<carveout_bytes> stage_count) {
|
||||
// For Mxf8f6f4 sub-bytes, ElementA/B will be passed in as uint8_t
|
||||
// Each stage include (CollectiveMma::SharedStorage)
|
||||
// 1. smem for A and smem for B (CollectiveMma::SharedStorage::TensorStorage)
|
||||
// 2. one MainloopPipeline = PipelineTmaUmmaAsync (CollectiveMma::SharedStorage::SharedStorage)
|
||||
// 3. smem for SFB and smem for SFB (CollectiveMma::SharedStorage::TensorStorage, independent of input size b.c. sizeof(sf) is fixed)
|
||||
constexpr auto mainloop_pipeline_bytes = sizeof(typename cutlass::PipelineTmaUmmaAsync<1>::SharedStorage);
|
||||
constexpr auto a_bits = cute::sizeof_bits_v<ElementA>;
|
||||
constexpr auto b_bits = cute::sizeof_bits_v<ElementB>;
|
||||
constexpr auto stage_sfa_bytes = size(filter_zeros(TileShapeSFA{}));
|
||||
constexpr auto stage_sfb_bytes = size(filter_zeros(TileShapeSFB{}));
|
||||
|
||||
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 + stage_sfa_bytes + stage_sfb_bytes);
|
||||
|
||||
return (CapacityBytes - carveout_bytes) / stage_bytes;
|
||||
}
|
||||
|
||||
template <class ClusterShapeMNK, class AtomThrId>
|
||||
constexpr auto
|
||||
sm100_cluster_shape_to_tma_atom_SFB(ClusterShapeMNK cluster_shape_mnk, AtomThrId atom_thr_id) {
|
||||
static_assert(cute::rank(cluster_shape_mnk) == 3);
|
||||
if constexpr (cute::size(atom_thr_id) == 2) {
|
||||
// Always could use multicast feature for SFB with 2cta MMA.
|
||||
return cute::SM100_TMA_2SM_LOAD_MULTICAST{};
|
||||
}
|
||||
else if constexpr (size(atom_thr_id) == 1) {
|
||||
return detail::sm90_cluster_shape_to_tma_atom(cute::size<0>(cluster_shape_mnk));
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<ClusterShapeMNK>,
|
||||
"Unsupported Configuration for SM100 TMA");
|
||||
}
|
||||
}
|
||||
|
||||
namespace blockscaled {
|
||||
|
||||
enum class BlockScaledInstr {
|
||||
MXF4_NVF4,
|
||||
MXF4F6F8
|
||||
};
|
||||
|
||||
template <class KernelScheduleType, class T>
|
||||
struct blockscaled_type {};
|
||||
|
||||
template <class KernelScheduleType, class T, class SF>
|
||||
struct blockscaled_type<KernelScheduleType, cute::tuple<T,SF>> {
|
||||
using sf_type = SF;
|
||||
using data_type = T;
|
||||
static constexpr uint32_t SfVectorSize = detail::find_vector_size<KernelScheduleType>();
|
||||
};
|
||||
|
||||
template <class KernelScheduleType, class T, class SF, int SfVectorSize_>
|
||||
struct blockscaled_type<KernelScheduleType, cute::tuple<T,SF, cute::Int<SfVectorSize_>>> {
|
||||
using sf_type = SF;
|
||||
using data_type = T;
|
||||
static constexpr uint32_t SfVectorSize = SfVectorSize_;
|
||||
};
|
||||
|
||||
template <class KernelScheduleType, class T>
|
||||
struct blockscaled_type<KernelScheduleType, cutlass::mx_float6_t<T>> {
|
||||
using sf_type = cutlass::float_ue8m0_t;
|
||||
using data_type = T;
|
||||
static constexpr uint32_t SfVectorSize = 32;
|
||||
};
|
||||
|
||||
template <class KernelScheduleType, class T>
|
||||
struct blockscaled_type<KernelScheduleType, cutlass::mx_float4_t<T>> {
|
||||
using sf_type = cutlass::float_ue8m0_t;
|
||||
using data_type = T;
|
||||
static constexpr uint32_t SfVectorSize = 32;
|
||||
};
|
||||
|
||||
template <class KernelScheduleType, class T>
|
||||
struct blockscaled_type<KernelScheduleType, nv_float4_t<T>> {
|
||||
using sf_type = cutlass::float_ue4m3_t;
|
||||
using data_type = T;
|
||||
static constexpr uint32_t SfVectorSize = 16;
|
||||
};
|
||||
template <class KernelScheduleType, class T>
|
||||
struct blockscaled_type<KernelScheduleType, cutlass::mx_float8_t<T>> {
|
||||
using sf_type = cutlass::float_ue8m0_t;
|
||||
using data_type = T;
|
||||
static constexpr uint32_t SfVectorSize = 32;
|
||||
};
|
||||
|
||||
template <
|
||||
class KernelScheduleType,
|
||||
class ElementPairA, class ElementPairB,
|
||||
UMMA::Major UmmaMajorA, UMMA::Major UmmaMajorB
|
||||
>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static constexpr bool
|
||||
check_input_datatypes() {
|
||||
using ElementSFA = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::sf_type;
|
||||
using ElementSFB = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::sf_type;
|
||||
using ElementA = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::data_type;
|
||||
using ElementB = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::data_type;
|
||||
constexpr uint32_t SfVectorSizeA = detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::SfVectorSize;
|
||||
constexpr uint32_t SfVectorSizeB = detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::SfVectorSize;
|
||||
|
||||
auto is_auto_instr_selection_policy = [&]() {
|
||||
return ((cute::is_same_v<KernelScheduleType, KernelScheduleAuto>) ||
|
||||
(cute::is_same_v<KernelScheduleType, KernelScheduleBlockScaledGemmSm100>) ||
|
||||
(cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecialized1SmBlockScaledSm100>) ||
|
||||
(cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecialized2SmBlockScaledSm100>) ||
|
||||
(cute::is_same_v<KernelSchedulePtrArrayBlockScaledGemmSm100, KernelScheduleType>) ||
|
||||
(cute::is_same_v<KernelPtrArrayTmaWarpSpecialized1SmBlockScaledSm100, KernelScheduleType>) ||
|
||||
(cute::is_same_v<KernelPtrArrayTmaWarpSpecialized2SmBlockScaledSm100, KernelScheduleType>));
|
||||
};
|
||||
|
||||
static_assert(cute::is_same_v<ElementSFA, ElementSFB>, "Scale factor types for A and B should be the same.");
|
||||
static_assert((SfVectorSizeA == SfVectorSizeB), "Scale factor vector size for A and B should be the same.");
|
||||
if constexpr ((SfVectorSizeA == 0) || (SfVectorSizeB == 0)) {
|
||||
static_assert(!is_auto_instr_selection_policy(), "Auto instr selection isn't valid if scale factor vector size can't be determined from the types");
|
||||
}
|
||||
|
||||
static_assert(cute::is_same_v<ElementSFA, cutlass::float_ue8m0_t>
|
||||
|| cute::is_same_v<ElementSFA, cutlass::float_ue4m3_t>, "Incorrect scale factor type");
|
||||
|
||||
if constexpr (((sizeof_bits_v<ElementA> == 4 || sizeof_bits_v<ElementA> == 6 || sizeof_bits_v<ElementA> == 8) &&
|
||||
(sizeof_bits_v<ElementB> == 4 || sizeof_bits_v<ElementB> == 6 || sizeof_bits_v<ElementB> == 8) ) && // A and B are 4, 6, or 8 bit types and
|
||||
(!(sizeof_bits_v<ElementA> == 4 && sizeof_bits_v<ElementB> == 4) ) // A and B are not both 4 bit types
|
||||
) {
|
||||
///////////////////////////////////////////////////////////////////////
|
||||
// Mixed Precision FP4, FP6, FP8 case. -> MX_F4F6F8 instructions
|
||||
///////////////////////////////////////////////////////////////////////
|
||||
// 1. Check Scale factor data type
|
||||
static_assert(cute::is_same_v<ElementSFA, cutlass::float_ue8m0_t>, "MX_F4F6F8 only supports ue8m0 SF type");
|
||||
// 2. Check whether A and B type combinations are valid or not
|
||||
static_assert(
|
||||
( // If runtime datatypes are used, then both A and B should be runtime data type
|
||||
(
|
||||
cute::is_same_v<ElementA, cutlass::type_erased_dynamic_float8_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::type_erased_dynamic_float6_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::type_erased_dynamic_float4_t>
|
||||
) &&
|
||||
(
|
||||
cute::is_same_v<ElementB, cutlass::type_erased_dynamic_float8_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::type_erased_dynamic_float6_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::type_erased_dynamic_float4_t>
|
||||
)
|
||||
) ||
|
||||
( // Valid (explicit) A and B type pairs
|
||||
(
|
||||
cute::is_same_v<ElementA, cutlass::float_e2m1_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::float_e2m3_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::float_e3m2_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::float_e4m3_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::float_e5m2_t>
|
||||
) &&
|
||||
(
|
||||
cute::is_same_v<ElementB, cutlass::float_e2m1_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::float_e2m3_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::float_e3m2_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::float_e4m3_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::float_e5m2_t>
|
||||
)
|
||||
), "Incorrect types for A and B for MX_F4F6F8"
|
||||
);
|
||||
// 3. Check Scale factor vector size is valid.
|
||||
// Only SfVectorSize = 32 is allowed.
|
||||
static_assert((SfVectorSizeA == 32) && (SfVectorSizeB == 32), "Incorrect SfVectorSize for MX_F4F6F8 is deduced. SfVectorSize should be 32.");
|
||||
// 4. Check the kernel policy. Kernel policy should be either auto or *MXf8f6f4*
|
||||
static_assert((cute::is_base_of_v<KernelScheduleMxf8f6f4Sm100, KernelScheduleType> ||
|
||||
cute::is_base_of_v<KernelSchedulePtrArrayMxf8f6f4Sm100, KernelScheduleType> ||
|
||||
is_auto_instr_selection_policy()), "Incorrect Kernel Schedule Policy for Mx_F4F6F8 type inputs.");
|
||||
|
||||
return true;
|
||||
}
|
||||
else if constexpr ((sizeof_bits_v<ElementA> == 4 && sizeof_bits_v<ElementB> == 4)) {
|
||||
///////////////////////////////////////////////////////////////////////
|
||||
// A and B are both 4 bit types
|
||||
// There are multiple block scaled tcgen05.mma instructions supporting F4 types.
|
||||
///////////////////////////////////////////////////////////////////////
|
||||
|
||||
// 1. Check Scale factor data type
|
||||
static_assert(cute::is_same_v<ElementSFA, cutlass::float_ue8m0_t>
|
||||
|| cute::is_same_v<ElementSFA, cutlass::float_ue4m3_t>
|
||||
, "MXNV_F4 supports ue8m0 and ue4m3 SF types");
|
||||
// 2. Check whether A and B type combinations are valid or not
|
||||
static_assert(
|
||||
( // If runtime datatypes are used, then both A and B should be runtime data type
|
||||
cute::is_same_v<ElementA, cutlass::type_erased_dynamic_float4_t> &&
|
||||
cute::is_same_v<ElementB, cutlass::type_erased_dynamic_float4_t>
|
||||
) ||
|
||||
( // Valid (explicit) A and B type pairs
|
||||
(
|
||||
cute::is_same_v<ElementA, cutlass::float_e2m1_t>
|
||||
) &&
|
||||
(
|
||||
cute::is_same_v<ElementB, cutlass::float_e2m1_t>
|
||||
)
|
||||
), "Incorrect types for A and B for MXNV_F4");
|
||||
// 3. Skip checking the scale factor vector size. Will be checked later for specific Kernel Schedule policies.
|
||||
// 4. Check the kernel policy.
|
||||
static_assert((cute::is_base_of_v<KernelScheduleMxf8f6f4Sm100, KernelScheduleType> ||
|
||||
cute::is_base_of_v<KernelSchedulePtrArrayMxf8f6f4Sm100, KernelScheduleType> ||
|
||||
cute::is_base_of_v<KernelScheduleMxNvf4Sm100, KernelScheduleType> ||
|
||||
cute::is_base_of_v<KernelSchedulePtrArrayMxNvf4Sm100, KernelScheduleType> ||
|
||||
is_auto_instr_selection_policy()), "Incorrect Kernel Schedule Policy for F4 type inputs.");
|
||||
// If a policy is specified, do more checks
|
||||
if constexpr (cute::is_base_of_v<KernelScheduleMxf8f6f4Sm100, KernelScheduleType>
|
||||
|| cute::is_base_of_v<KernelSchedulePtrArrayMxf8f6f4Sm100, KernelScheduleType>
|
||||
) {
|
||||
// Perform additional checks. Only subset of FP4 and scale factor types are supported.
|
||||
static_assert(cute::is_same_v<ElementSFA, cutlass::float_ue8m0_t>, "MX_F4F6F8 only supports ue8m0 SF type");
|
||||
static_assert((cute::is_same_v<ElementA, cutlass::type_erased_dynamic_float4_t> && cute::is_same_v<ElementB, cutlass::type_erased_dynamic_float4_t>) ||
|
||||
(cute::is_same_v<ElementA, cutlass::float_e2m1_t> && cute::is_same_v<ElementB, cutlass::float_e2m1_t>), "Incorrect types for A and B for MX_F4F6F8");
|
||||
static_assert((SfVectorSizeA == 32) && (SfVectorSizeB == 32), "Incorrect SfVectorSize for MX_F4F6F8 is deduced. SfVectorSize should be 32.");
|
||||
return true;
|
||||
}
|
||||
else if constexpr (cute::is_base_of_v<KernelScheduleMxNvf4Sm100, KernelScheduleType>
|
||||
|| cute::is_base_of_v<KernelSchedulePtrArrayMxNvf4Sm100, KernelScheduleType>
|
||||
) {
|
||||
static_assert((UmmaMajorA == UMMA::Major::K && UmmaMajorB == UMMA::Major::K), "MX/NV_F4 only supports RowMajor A, and ColMajorB");
|
||||
static_assert(detail::find_vector_size<KernelScheduleType>() == SfVectorSizeA, "Kernel Schedule policy doesn't match the scale factor vector size.");
|
||||
return true;
|
||||
}
|
||||
else { // auto policy
|
||||
// If the scale factor type is ue4m3 or the scale factor vector size is 16 -> only MXF4_NVF4 instruction can support it
|
||||
// For MXF4_NVF4, the layouts should be RowMajor A, and ColMajorB
|
||||
static_assert(is_auto_instr_selection_policy(), "Kernel Schedule policy should be auto");
|
||||
if constexpr (SfVectorSizeA == 16 || SfVectorSizeB == 16
|
||||
|| cute::is_same_v<ElementSFA, cutlass::float_ue4m3_t>
|
||||
) { // Only MXF4NVF4 can support these types
|
||||
static_assert((UmmaMajorA == UMMA::Major::K && UmmaMajorB == UMMA::Major::K), "NV_F4 only supports RowMajor A, and ColMajorB");
|
||||
return true;
|
||||
}
|
||||
return true;
|
||||
}
|
||||
}
|
||||
else {
|
||||
return false;
|
||||
}
|
||||
return false;
|
||||
}
|
||||
|
||||
template <
|
||||
class TileShape_MNK, // (MmaAtomShape_M, MmaAtomShape_N, CtaTileShapeK)
|
||||
class ClusterShape_MNK,
|
||||
class KernelScheduleType
|
||||
>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static constexpr bool
|
||||
is_2sm() {
|
||||
// 2SM kernel schedule is requested
|
||||
if constexpr (cute::is_base_of_v<KernelSchedule2Sm, KernelScheduleType>) { return true; }
|
||||
// 1SM kernel schedule is requested
|
||||
else if constexpr (cute::is_base_of_v<KernelSchedule1Sm, KernelScheduleType>) { return false; }
|
||||
// auto schedule is used.
|
||||
else {
|
||||
if constexpr (!cute::is_static_v<ClusterShape_MNK>) {
|
||||
// If the cluster shape is dynamic, we can't guarantee 2x1. Default to 1sm.
|
||||
// If tile shape M is 256, throw an error. M=256 is only supported by 2SM instructions.
|
||||
static_assert(get<0>(TileShape_MNK{}) != 256, "If M=256, auto policy can't create 2sm kernels. Specify a 2SM policy");
|
||||
return false;
|
||||
}
|
||||
else if constexpr (cute::is_static_v<ClusterShape_MNK> && cute::get<0>(ClusterShape_MNK{}) % 2 == 0) {
|
||||
// We need to check the TileShape
|
||||
if constexpr (get<0>(TileShape_MNK{}) == 256) {
|
||||
return true;
|
||||
}
|
||||
else if constexpr (get<0>(TileShape_MNK{}) == 128) {
|
||||
return false;
|
||||
}
|
||||
else {
|
||||
static_assert(get<0>(TileShape_MNK{}) == 0, "Unsupported M dimension for TileShape_MNK.");
|
||||
}
|
||||
}
|
||||
else { return false;}
|
||||
}
|
||||
}
|
||||
|
||||
template <
|
||||
class ElementPairA,
|
||||
class ElementPairB,
|
||||
class ElementAccumulator,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
class KernelScheduleType
|
||||
>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static constexpr auto
|
||||
select_instr() {
|
||||
using ElementSFA = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::sf_type;
|
||||
using ElementSFB = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::sf_type;
|
||||
using ElementA = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::data_type;
|
||||
using ElementB = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::data_type;
|
||||
constexpr uint32_t SfVectorSizeA = detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::SfVectorSize;
|
||||
constexpr uint32_t SfVectorSizeB = detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::SfVectorSize;
|
||||
constexpr int SFVectorSize = SfVectorSizeA > SfVectorSizeB ? SfVectorSizeA : SfVectorSizeB;
|
||||
using ElementSF = ElementSFA;
|
||||
|
||||
if constexpr (cute::is_base_of_v<KernelScheduleMxf8f6f4Sm100, KernelScheduleType>
|
||||
|| cute::is_base_of_v<KernelSchedulePtrArrayMxf8f6f4Sm100, KernelScheduleType>
|
||||
) {
|
||||
return detail::blockscaled::BlockScaledInstr::MXF4F6F8;
|
||||
}
|
||||
else if constexpr (cute::is_base_of_v<KernelScheduleMxNvf4Sm100, KernelScheduleType>
|
||||
|| cute::is_base_of_v<KernelSchedulePtrArrayMxNvf4Sm100, KernelScheduleType>
|
||||
) {
|
||||
return detail::blockscaled::BlockScaledInstr::MXF4_NVF4;
|
||||
}
|
||||
else {
|
||||
// Auto scheduling
|
||||
if constexpr ((sizeof_bits_v<ElementA> >= 6 && sizeof_bits_v<ElementA> <= 8) &&
|
||||
(sizeof_bits_v<ElementB> >= 6 && sizeof_bits_v<ElementB> <= 8)) {
|
||||
// These types can only be supported by MX_F8F6F4 instruction
|
||||
static_assert(SFVectorSize == 32, "Incorrect SF vector size");
|
||||
return detail::blockscaled::BlockScaledInstr::MXF4F6F8;
|
||||
}
|
||||
else if constexpr (( sizeof_bits_v<ElementA> == 4 && (sizeof_bits_v<ElementB> == 6 || sizeof_bits_v<ElementB> == 8)) ||
|
||||
((sizeof_bits_v<ElementA> == 6 || sizeof_bits_v<ElementA> == 8) && sizeof_bits_v<ElementB> == 4)) {
|
||||
// Fp4 can be mixed with FP6, Fp8 with Mxf8f6f4 only
|
||||
return detail::blockscaled::BlockScaledInstr::MXF4F6F8;
|
||||
}
|
||||
else if constexpr (sizeof_bits_v<ElementA> == 4 && sizeof_bits_v<ElementB> == 4) {
|
||||
// Both A and B are 4bits
|
||||
if constexpr (UmmaMajorA == UMMA::Major::K && UmmaMajorB == UMMA::Major::K) {
|
||||
// MXF4_NVF4 possible
|
||||
return detail::blockscaled::BlockScaledInstr::MXF4_NVF4;
|
||||
}
|
||||
else {
|
||||
static_assert(SFVectorSize == 32, "Incorrect SF vector size");
|
||||
static_assert( cute::is_same_v<ElementSF, cutlass::float_ue8m0_t> &&
|
||||
(cute::is_same_v<ElementA, cutlass::float_e2m1_t> && cute::is_same_v<ElementB, cutlass::float_e2m1_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::type_erased_dynamic_float4_t> && cute::is_same_v<ElementB, cutlass::type_erased_dynamic_float4_t>),
|
||||
"Only MXF4 support with non-TN and Mxf8f6f4");
|
||||
return detail::blockscaled::BlockScaledInstr::MXF4F6F8;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace blockscaled
|
||||
|
||||
template <
|
||||
class ElementPairA,
|
||||
class ElementPairB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
detail::blockscaled::BlockScaledInstr Instr,
|
||||
class KernelScheduleType
|
||||
>
|
||||
constexpr auto
|
||||
sm100_make_blockscaled_1sm_trivial_tiled_mma() {
|
||||
// For MMA_1sm atoms, the MMA's AtomLayout is same as the ClusterShape
|
||||
using AtomLayout_MNK = Layout<ClusterShape_MNK>;
|
||||
constexpr int M = cute::size<0>(TileShape_MNK{});
|
||||
static_assert(M == 128, "Invalid TileShape_M.");
|
||||
|
||||
// Do not allow a tiled MMA N mode > 1, as that is not reasonable.
|
||||
constexpr int N = cute::size<1>(TileShape_MNK{});
|
||||
static_assert(N == 128 || N == 192 || N == 256, "Invalid TileShape_N.");
|
||||
|
||||
using ElementSFA = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::sf_type;
|
||||
using ElementSFB = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::sf_type;
|
||||
using ElementA = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::data_type;
|
||||
using ElementB = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::data_type;
|
||||
constexpr uint32_t SfVectorSizeA = detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::SfVectorSize;
|
||||
[[maybe_unused]] constexpr uint32_t SfVectorSizeB = detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::SfVectorSize;
|
||||
|
||||
using ElementAMma = decltype(cutlass::gemm::collective::detail::sm100_kernel_input_element_to_mma_input_element<ElementA, Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8>());
|
||||
using ElementBMma = decltype(cutlass::gemm::collective::detail::sm100_kernel_input_element_to_mma_input_element<ElementB, Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8>());
|
||||
|
||||
using ElementSF = ElementSFA;
|
||||
|
||||
if constexpr (Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8) {
|
||||
|
||||
return make_tiled_mma(cute::SM100_MMA_MXF8F6F4_SS<ElementAMma, ElementBMma, ElementAccumulator, ElementSF,
|
||||
M, N, UmmaMajorA, UmmaMajorB>{});
|
||||
}
|
||||
else if constexpr (Instr == detail::blockscaled::BlockScaledInstr::MXF4_NVF4) {
|
||||
constexpr int SFVectorSize = SfVectorSizeA;
|
||||
return make_tiled_mma(cute::SM100_MMA_MXF4_SS<ElementAMma, ElementBMma, ElementAccumulator, ElementSF,
|
||||
M, N, SFVectorSize, UmmaMajorA, UmmaMajorB>{});
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<ElementAMma>,
|
||||
"Unsupported configuration for SM100 collective builder.");
|
||||
}
|
||||
}
|
||||
|
||||
template <
|
||||
class ElementPairA,
|
||||
class ElementPairB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
detail::blockscaled::BlockScaledInstr Instr,
|
||||
class KernelScheduleType
|
||||
>
|
||||
constexpr auto
|
||||
sm100_make_blockscaled_2sm_trivial_tiled_mma() {
|
||||
|
||||
constexpr int M = cute::size<0>(TileShape_MNK{});
|
||||
static_assert(M == 256, "Invalid TileShape_M.");
|
||||
|
||||
// Do not allow a tiled MMA N mode > 1, as that is not reasonable.
|
||||
constexpr int N = cute::size<1>(TileShape_MNK{});
|
||||
static_assert(N == 128 || N == 192 || N == 256, "Invalid TileShape_N.");
|
||||
|
||||
using ElementSFA = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::sf_type;
|
||||
using ElementSFB = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::sf_type;
|
||||
using ElementA = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::data_type;
|
||||
using ElementB = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::data_type;
|
||||
constexpr uint32_t SfVectorSizeA = detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::SfVectorSize;
|
||||
[[maybe_unused]] constexpr uint32_t SfVectorSizeB = detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::SfVectorSize;
|
||||
using ElementAMma = decltype(cutlass::gemm::collective::detail::sm100_kernel_input_element_to_mma_input_element<ElementA, Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8>());
|
||||
using ElementBMma = decltype(cutlass::gemm::collective::detail::sm100_kernel_input_element_to_mma_input_element<ElementB, Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8>());
|
||||
|
||||
using ElementSF = ElementSFA;
|
||||
|
||||
if constexpr (Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8) {
|
||||
return make_tiled_mma(cute::SM100_MMA_MXF8F6F4_2x1SM_SS<ElementAMma, ElementBMma, ElementAccumulator, ElementSF,
|
||||
M, N, UmmaMajorA, UmmaMajorB>{});
|
||||
}
|
||||
else if constexpr (Instr == detail::blockscaled::BlockScaledInstr::MXF4_NVF4) {
|
||||
constexpr int SFVectorSize = SfVectorSizeA > SfVectorSizeB ? SfVectorSizeA : SfVectorSizeB;
|
||||
return make_tiled_mma(cute::SM100_MMA_MXF4_2x1SM_SS<ElementAMma, ElementBMma, ElementAccumulator, ElementSF,
|
||||
M, N, SFVectorSize, UmmaMajorA, UmmaMajorB>{});
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<ElementAMma>,
|
||||
"Unsupported configuration for SM100 collective builder.");
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
template <
|
||||
class ElementPairA,
|
||||
class ElementPairB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
detail::blockscaled::BlockScaledInstr Instr,
|
||||
class KernelScheduleType,
|
||||
bool Is2SM
|
||||
>
|
||||
struct TrivialBlockscaledMma {};
|
||||
|
||||
template <
|
||||
class ElementPairA,
|
||||
class ElementPairB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
detail::blockscaled::BlockScaledInstr Instr,
|
||||
class KernelScheduleType
|
||||
>
|
||||
struct TrivialBlockscaledMma <
|
||||
ElementPairA,
|
||||
ElementPairB,
|
||||
ElementAccumulator,
|
||||
TileShape_MNK,
|
||||
ClusterShape_MNK,
|
||||
UmmaMajorA,
|
||||
UmmaMajorB,
|
||||
Instr,
|
||||
KernelScheduleType,
|
||||
true /*Is2SM*/> {
|
||||
using type = decltype(sm100_make_blockscaled_2sm_trivial_tiled_mma<ElementPairA, ElementPairB, ElementAccumulator,
|
||||
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, Instr, KernelScheduleType>());
|
||||
};
|
||||
|
||||
template <
|
||||
class ElementPairA,
|
||||
class ElementPairB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
detail::blockscaled::BlockScaledInstr Instr,
|
||||
class KernelScheduleType
|
||||
>
|
||||
struct TrivialBlockscaledMma<
|
||||
ElementPairA,
|
||||
ElementPairB,
|
||||
ElementAccumulator,
|
||||
TileShape_MNK,
|
||||
ClusterShape_MNK,
|
||||
UmmaMajorA,
|
||||
UmmaMajorB,
|
||||
Instr,
|
||||
KernelScheduleType,
|
||||
false /*Is2SM*/> {
|
||||
using type = decltype(sm100_make_blockscaled_1sm_trivial_tiled_mma<ElementPairA, ElementPairB, ElementAccumulator,
|
||||
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, Instr, KernelScheduleType>());
|
||||
};
|
||||
} // namespace detail
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
class ElementPairA,
|
||||
class GmemLayoutATag,
|
||||
int AlignmentA,
|
||||
class ElementPairB,
|
||||
class GmemLayoutBTag,
|
||||
int AlignmentB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK, // (MmaAtomShapeM, MmaAtomShapeN, TileK)
|
||||
class ClusterShape_MNK, // Static cluster shape or dynamic (int, int, _1)
|
||||
class StageCountType,
|
||||
class KernelScheduleType
|
||||
>
|
||||
struct CollectiveBuilder<
|
||||
arch::Sm100,
|
||||
arch::OpClassBlockScaledTensorOp,
|
||||
ElementPairA,
|
||||
GmemLayoutATag,
|
||||
AlignmentA,
|
||||
ElementPairB,
|
||||
GmemLayoutBTag,
|
||||
AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape_MNK,
|
||||
ClusterShape_MNK,
|
||||
StageCountType,
|
||||
KernelScheduleType,
|
||||
cute::enable_if_t<
|
||||
// Blockscaled Gemm
|
||||
(cute::is_base_of_v<KernelScheduleBlockScaledGemmSm100, KernelScheduleType> ||
|
||||
cute::is_same_v<KernelScheduleAuto, KernelScheduleType>)
|
||||
&&
|
||||
// Alignment check
|
||||
detail::sm1xx_blockscaled_gemm_is_aligned<typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::data_type,
|
||||
AlignmentA,
|
||||
typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::data_type,
|
||||
AlignmentB,
|
||||
KernelScheduleType>()>>
|
||||
{
|
||||
using ElementSFA = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::sf_type;
|
||||
using ElementSFB = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::sf_type;
|
||||
using ElementA = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairA>::data_type;
|
||||
using ElementB = typename detail::blockscaled::blockscaled_type<KernelScheduleType, ElementPairB>::data_type;
|
||||
using ElementSF = ElementSFA;
|
||||
|
||||
static constexpr cute::UMMA::Major UmmaMajorA = cutlass::gemm::collective::detail::tag_to_umma_major_A<GmemLayoutATag>();
|
||||
static constexpr cute::UMMA::Major UmmaMajorB = cutlass::gemm::collective::detail::tag_to_umma_major_B<GmemLayoutBTag>();
|
||||
|
||||
static_assert(cute::is_static_v<TileShape_MNK>, "TileShape has to be static");
|
||||
static_assert(detail::blockscaled::check_input_datatypes<KernelScheduleType, ElementPairA, ElementPairB, UmmaMajorA, UmmaMajorB>(), "Incorrect input types");
|
||||
|
||||
static constexpr bool is_2sm = detail::blockscaled::is_2sm<TileShape_MNK, ClusterShape_MNK, KernelScheduleType>();
|
||||
static constexpr auto Instr = detail::blockscaled::select_instr<ElementPairA, ElementPairB, ElementAccumulator, UmmaMajorA, UmmaMajorB, KernelScheduleType>();
|
||||
|
||||
using TiledMma = typename cutlass::gemm::collective::detail::TrivialBlockscaledMma<ElementPairA, ElementPairB, ElementAccumulator,
|
||||
TileShape_MNK, ClusterShape_MNK,
|
||||
UmmaMajorA, UmmaMajorB, Instr, KernelScheduleType, is_2sm>::type;
|
||||
|
||||
static constexpr bool UseMxf8f6f4 = Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8;
|
||||
|
||||
static_assert(UseMxf8f6f4 || (cutlass::gemm::detail::is_k_major_A<GmemLayoutATag>() && cutlass::gemm::detail::is_k_major_B<GmemLayoutBTag>()), "Only Mxf8f6f4 supports non-K major inputs");
|
||||
|
||||
// Data type used by MMA instruction
|
||||
using ElementAMma = decltype(cutlass::gemm::collective::detail::sm100_kernel_input_element_to_mma_input_element<ElementA, UseMxf8f6f4>());
|
||||
using ElementBMma = decltype(cutlass::gemm::collective::detail::sm100_kernel_input_element_to_mma_input_element<ElementB, UseMxf8f6f4>());
|
||||
|
||||
static_assert(detail::sm100_gemm_check_for_f8f6f4_mix8bit_requirement<ElementAMma, ElementBMma,
|
||||
TileShape_MNK, ClusterShape_MNK,
|
||||
UmmaMajorA, UmmaMajorB, KernelScheduleType, is_2sm>(),
|
||||
"TileSize and MNK Major does not met with MMA Mix 8-bit TMA load requirement" );
|
||||
|
||||
static constexpr uint32_t SFVectorSize = TiledMma::SFVecSize;
|
||||
|
||||
// Basic storage block for new Scaling Factor Layouts
|
||||
using AtomThrID = typename TiledMma::AtomThrID;
|
||||
using Sm100BlkScaledConfig = cutlass::detail::Sm100BlockScaledConfig<SFVectorSize>;
|
||||
|
||||
using ElementAMma_SmemAllocType = cute::conditional_t<UseMxf8f6f4, uint8_t, ElementAMma>;
|
||||
using ElementBMma_SmemAllocType = cute::conditional_t<UseMxf8f6f4, uint8_t, ElementBMma>;
|
||||
|
||||
// ((MMA_TILE_M,MMA_TILE_K), MMA_M, MMA_K)
|
||||
using MmaShapeA_MK = decltype(partition_shape_A(TiledMma{}, make_shape(cute::size<0>(TileShape_MNK{}),
|
||||
cute::size<2>(TileShape_MNK{}))));
|
||||
// ((MMA_TILE_N,MMA_TILE_K), MMA_N, MMA_K)
|
||||
using MmaShapeB_NK = decltype(partition_shape_B(TiledMma{}, make_shape(cute::size<1>(TileShape_MNK{}),
|
||||
cute::size<2>(TileShape_MNK{}))));
|
||||
|
||||
using GmemTiledCopyA = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_A(
|
||||
ClusterShape_MNK{}, AtomThrID{}));
|
||||
|
||||
using GmemTiledCopyB = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_B(
|
||||
ClusterShape_MNK{}, AtomThrID{}));
|
||||
|
||||
using GmemTiledCopySFA = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_A(
|
||||
ClusterShape_MNK{}, AtomThrID{}));
|
||||
|
||||
using GmemTiledCopySFB = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_SFB(
|
||||
ClusterShape_MNK{}, AtomThrID{}));
|
||||
|
||||
using GmemTiledCopyPairA = decltype(cute::make_tuple(GmemTiledCopyA{}, GmemTiledCopySFA{}));
|
||||
using GmemTiledCopyPairB = decltype(cute::make_tuple(GmemTiledCopyB{}, GmemTiledCopySFB{}));
|
||||
|
||||
//
|
||||
// Construct SMEM layout (SmemLayoutAtom) for A and SFA
|
||||
//
|
||||
using BlockTileA_M = decltype(cute::size<0,0>(MmaShapeA_MK{}) * cute::size<1>(MmaShapeA_MK{}));
|
||||
using BlockTileA_K = decltype(cute::size<0,1>(MmaShapeA_MK{}) * cute::size<2>(MmaShapeA_MK{}));
|
||||
using SmemLayoutAtomA = decltype(cutlass::gemm::collective::detail::sm100_smem_selector<
|
||||
UmmaMajorA, ElementAMma_SmemAllocType, BlockTileA_M, BlockTileA_K>());
|
||||
|
||||
// A single indivisible block will hold 4 scale factors of 128 rows/columns (A/B matrix).
|
||||
// 4 is chosen to make consecutive 32bits of data to have scale factors for only a single row (col). 32bits corresponds to the TMEM word size
|
||||
using Blk_MN = typename Sm100BlkScaledConfig::Blk_MN;
|
||||
using Blk_SF = typename Sm100BlkScaledConfig::Blk_SF;
|
||||
using Blk_Elems = decltype(Blk_MN{} * Blk_SF{});
|
||||
using SmemLayoutAtomSFA = decltype(Sm100BlkScaledConfig::deduce_smem_layoutSFA(TiledMma{}, TileShape_MNK{}));
|
||||
using SmemLayoutAtomsA = decltype(cute::make_tuple(SmemLayoutAtomA{}, SmemLayoutAtomSFA{}));
|
||||
|
||||
//
|
||||
// Construct SMEM layout (SmemLayoutAtom) for B and SFB
|
||||
//
|
||||
using BlockTileB_N = decltype(cute::size<0,0>(MmaShapeB_NK{}) * cute::size<1>(MmaShapeB_NK{}));
|
||||
using BlockTileB_K = decltype(cute::size<0,1>(MmaShapeB_NK{}) * cute::size<2>(MmaShapeB_NK{}));
|
||||
using SmemLayoutAtomB = decltype(cutlass::gemm::collective::detail::sm100_smem_selector<
|
||||
UmmaMajorB, ElementBMma_SmemAllocType, BlockTileB_N, BlockTileB_K>());
|
||||
using SmemLayoutAtomSFB = decltype(Sm100BlkScaledConfig::deduce_smem_layoutSFB(TiledMma{}, TileShape_MNK{}));
|
||||
using SmemLayoutAtomsB = decltype(cute::make_tuple(SmemLayoutAtomB{}, SmemLayoutAtomSFB{}));
|
||||
|
||||
//
|
||||
// Construct Strides for A, SFA, B, and SFB
|
||||
//
|
||||
using StrideA = cutlass::gemm::TagToStrideA_t<GmemLayoutATag>;
|
||||
using StrideB = cutlass::gemm::TagToStrideB_t<GmemLayoutBTag>;
|
||||
using InternalStrideA = cute::remove_pointer_t<StrideA>;
|
||||
using InternalStrideB = cute::remove_pointer_t<StrideB>;
|
||||
using InternalLayoutSFA = decltype(Sm100BlkScaledConfig::deduce_layoutSFA());
|
||||
using InternalLayoutSFB = decltype(Sm100BlkScaledConfig::deduce_layoutSFB());
|
||||
using LayoutSFA = cute::conditional_t<cute::is_same_v<InternalStrideA, StrideA>, InternalLayoutSFA, InternalLayoutSFA *>;
|
||||
using LayoutSFB = cute::conditional_t<cute::is_same_v<InternalStrideB, StrideB>, InternalLayoutSFB, InternalLayoutSFB *>;
|
||||
using StridePairA = decltype(cute::make_tuple(StrideA{}, LayoutSFA{}));
|
||||
using StridePairB = decltype(cute::make_tuple(StrideB{}, LayoutSFB{}));
|
||||
|
||||
static constexpr int MMA_N = cute::size<1>(TileShape_MNK{});
|
||||
static constexpr uint32_t AccumulatorPipelineStageCount = (MMA_N == 256) ? 1 : 2;
|
||||
// Grouped GEMM (where Stride type is Stride*) does not use CLC based scheduler.
|
||||
static constexpr uint32_t SchedulerPipelineStageCount = cute::is_same_v<InternalStrideA, StrideA> ? 3 : 1;
|
||||
static constexpr bool IsArrayOfPointersGemm = cute::is_base_of_v<KernelSchedulePtrArrayBlockScaledGemmSm100, KernelScheduleType>;
|
||||
static constexpr uint32_t KernelSmemCarveout = detail::Sm100DenseGemmTmaUmmaCarveout<
|
||||
ClusterShape_MNK,
|
||||
AccumulatorPipelineStageCount,
|
||||
SchedulerPipelineStageCount,
|
||||
detail::CLCResponseSize,
|
||||
IsArrayOfPointersGemm,
|
||||
4 // 4 Tensor maps for A, SFA, B and SFB
|
||||
>::KernelSmemCarveout;
|
||||
// Reduce SMEM capacity available for buffers considering barrier allocations.
|
||||
static constexpr int Sm100ReducedSmemCapacityBytes = cutlass::gemm::collective::detail::sm100_smem_capacity_bytes - KernelSmemCarveout;
|
||||
|
||||
using SmemTileShape = cute::Shape<BlockTileA_M, BlockTileB_N, BlockTileA_K>;
|
||||
|
||||
static constexpr int PipelineStages = cutlass::gemm::collective::detail::sm100_compute_stage_count_or_override_blockscaled<
|
||||
Sm100ReducedSmemCapacityBytes, ElementAMma_SmemAllocType, ElementBMma_SmemAllocType, SmemTileShape, SmemLayoutAtomSFA, SmemLayoutAtomSFB>(StageCountType{});
|
||||
static_assert(PipelineStages > 0, "Smem usage is too high. Can't create any SMEM buffers for A, B, SFA, and SFB.");
|
||||
|
||||
using DispatchPolicy =
|
||||
cute::conditional_t<IsArrayOfPointersGemm,
|
||||
cutlass::gemm::MainloopSm100ArrayTmaUmmaWarpSpecializedBlockScaled<
|
||||
PipelineStages,
|
||||
SchedulerPipelineStageCount,
|
||||
AccumulatorPipelineStageCount,
|
||||
ClusterShape_MNK
|
||||
>,
|
||||
cutlass::gemm::MainloopSm100TmaUmmaWarpSpecializedBlockScaled<
|
||||
PipelineStages,
|
||||
SchedulerPipelineStageCount,
|
||||
AccumulatorPipelineStageCount,
|
||||
ClusterShape_MNK
|
||||
>
|
||||
>;
|
||||
|
||||
using CollectiveOp = cutlass::gemm::collective::CollectiveMma<
|
||||
DispatchPolicy,
|
||||
TileShape_MNK,
|
||||
cute::tuple<ElementA, ElementSF>,
|
||||
StridePairA,
|
||||
cute::tuple<ElementB, ElementSF>,
|
||||
StridePairB,
|
||||
TiledMma,
|
||||
GmemTiledCopyPairA,
|
||||
SmemLayoutAtomsA,
|
||||
void,
|
||||
cute::identity,
|
||||
GmemTiledCopyPairB,
|
||||
SmemLayoutAtomsB,
|
||||
void,
|
||||
cute::identity
|
||||
>;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,572 @@
|
||||
/***************************************************************************************************
|
||||
* 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/gemm.h"
|
||||
#include "cutlass/gemm/kernel/sm100_tile_scheduler.hpp"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp" // KernelSchedule1Sm, KernelSchedule2Sm
|
||||
#include "cutlass/gemm/collective/builders/sm90_common.inl" // detail::sm90_cluster_shape_to_tma_atom()
|
||||
#include "cutlass/numeric_types.h" // all numeric types
|
||||
#include "cutlass/detail/dependent_false.hpp" // detail::dependent_false
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/detail/layout.hpp" // cutlass::detail::get_input_alignment_bits()
|
||||
#include "cutlass/layout/matrix.h" // cutlass::layout::RowMajor, cutlass::layout::ColumnMajor
|
||||
#include "cutlass/fast_math.h" // cutlass::round_up, cutlass::const_max
|
||||
#include "cutlass/arch/arch.h"
|
||||
|
||||
#include "cute/atom/mma_traits_sm100.hpp" // UMMA::Layout_MN_SW*
|
||||
#include "cute/atom/copy_traits_sm100_tma.hpp" // SM100_TMA_*SM_LOAD_*
|
||||
#include "cute/arch/tmem_allocator_sm100.hpp"
|
||||
#include "cute/arch/mma_sm100_desc.hpp" // cute::UMMA::Major
|
||||
#include "cute/arch/mma_sm100_umma.hpp" // SM100_*MMA_SS_*
|
||||
#include "cute/numeric/integral_constant.hpp" // is_static_v, cute::integral_constant
|
||||
#include "cute/util/type_traits.hpp" // cute::alignment_of_v
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::collective {
|
||||
|
||||
// Forward Declaration
|
||||
struct KernelScheduleAuto;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
//
|
||||
// Some named constants
|
||||
//
|
||||
constexpr int sm100_smem_capacity_bytes = cutlass::arch::sm100_smem_capacity_bytes;
|
||||
constexpr int CLCResponseSize =
|
||||
sizeof(typename cutlass::gemm::kernel::detail::PersistentTileSchedulerSm100<Shape<_1,_1,_1>,1>::CLCResponse{});
|
||||
|
||||
// Maps input element to umma element
|
||||
template <class Element, bool IsF8F6F4 = true>
|
||||
constexpr auto
|
||||
sm100_kernel_input_element_to_mma_input_element() {
|
||||
if constexpr (cute::is_same_v<Element, float>) {
|
||||
return cutlass::tfloat32_t{};
|
||||
}
|
||||
else if constexpr (cute::is_same_v<Element, cutlass::float_e2m1_t> && IsF8F6F4) {
|
||||
return cutlass::detail::float_e2m1_unpacksmem_t{};
|
||||
}
|
||||
else if constexpr (cute::is_same_v<Element, cutlass::float_e3m2_t> && IsF8F6F4) {
|
||||
return cutlass::detail::float_e3m2_unpacksmem_t{};
|
||||
}
|
||||
else if constexpr (cute::is_same_v<Element, cutlass::float_e2m3_t> && IsF8F6F4) {
|
||||
return cutlass::detail::float_e2m3_unpacksmem_t{};
|
||||
}
|
||||
else if constexpr (cute::is_same_v<Element, cutlass::type_erased_dynamic_float4_t> && IsF8F6F4) {
|
||||
return cutlass::detail::type_erased_dynamic_float4_unpacksmem_t{};
|
||||
}
|
||||
else if constexpr (cute::is_same_v<Element, cutlass::type_erased_dynamic_float6_t> && IsF8F6F4) {
|
||||
return cutlass::detail::type_erased_dynamic_float6_unpacksmem_t{};
|
||||
}
|
||||
else {
|
||||
return Element{};
|
||||
}
|
||||
}
|
||||
|
||||
// Maps 2.x A matrix layout tag to respective UMMA major mode enum
|
||||
template <class Layout>
|
||||
constexpr cute::UMMA::Major
|
||||
tag_to_umma_major_A() {
|
||||
using LayoutA = cute::remove_pointer_t<Layout>;
|
||||
if constexpr (cute::is_same_v<LayoutA, cutlass::layout::RowMajor>) {
|
||||
return cute::UMMA::Major::K;
|
||||
}
|
||||
else if constexpr (cute::is_same_v<LayoutA, cutlass::layout::ColumnMajor>) {
|
||||
return cute::UMMA::Major::MN;
|
||||
}
|
||||
else if constexpr (cutlass::detail::is_major<0, LayoutA>()) {
|
||||
return cute::UMMA::Major::MN;
|
||||
}
|
||||
else if constexpr (cutlass::detail::is_major<1, LayoutA>()) {
|
||||
return cute::UMMA::Major::K;
|
||||
}
|
||||
else {
|
||||
static_assert(sizeof(LayoutA) == 0, "Invalid layout.");
|
||||
}
|
||||
}
|
||||
|
||||
// Maps 2.x B matrix layout tag to respective UMMA major mode enum
|
||||
template <class Layout>
|
||||
constexpr cute::UMMA::Major
|
||||
tag_to_umma_major_B() {
|
||||
using LayoutB = cute::remove_pointer_t<Layout>;
|
||||
if constexpr (cute::is_same_v<LayoutB, cutlass::layout::RowMajor>) {
|
||||
return cute::UMMA::Major::MN;
|
||||
}
|
||||
else if constexpr (cute::is_same_v<LayoutB, cutlass::layout::ColumnMajor>) {
|
||||
return cute::UMMA::Major::K;
|
||||
}
|
||||
else if constexpr (cutlass::detail::is_major<0, LayoutB>()) {
|
||||
return cute::UMMA::Major::MN;
|
||||
}
|
||||
else if constexpr (cutlass::detail::is_major<1, LayoutB>()) {
|
||||
return cute::UMMA::Major::K;
|
||||
}
|
||||
else {
|
||||
static_assert(sizeof(LayoutB) == 0, "Invalid layout.");
|
||||
}
|
||||
}
|
||||
|
||||
// Helper for SS UMMA smem selection that considers a tensor TileShape:
|
||||
// (BLK_MN, BLK_K)
|
||||
// or hierarchically
|
||||
// ((BLK_MN0,BLK_MN1,...),(BLK_K0,BLK_K1,...))
|
||||
// and returns the largest UMMA::Layout that fits BLK_MN0 and BLK_K0
|
||||
template <cute::UMMA::Major major, class ElementType, class BLK_MN, class BLK_K>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
sm100_smem_selector() {
|
||||
auto BLK_MN0 = size<0>(BLK_MN{});
|
||||
auto BLK_K0 = size<0>(BLK_K{});
|
||||
|
||||
static_assert(BLK_MN0 % 8 == 0, "BLK_MN0 must be a multiple of 8.");
|
||||
static_assert(BLK_K0 % 8 == 0, "BLK_K0 must be a multiple of 8.");
|
||||
|
||||
if constexpr (major == cute::UMMA::Major::MN) {
|
||||
// Handle the special case for F32 NT kernels
|
||||
if constexpr ((sizeof(ElementType) == 4)) {
|
||||
static_assert(BLK_MN0 % size<0>(UMMA::Layout_MN_SW128_32B_Atom<ElementType>{}) == 0, "for mn-major tf32 operands, SW128_32B is the only available smem layout");
|
||||
return UMMA::Layout_MN_SW128_32B_Atom<ElementType>{};
|
||||
}
|
||||
else {
|
||||
// All other data types are handled as SM90
|
||||
if constexpr (BLK_MN0 % size<0>(UMMA::Layout_MN_SW128_Atom<ElementType>{}) == 0) {
|
||||
return UMMA::Layout_MN_SW128_Atom<ElementType>{};
|
||||
}
|
||||
else if constexpr (BLK_MN0 % size<0>(UMMA::Layout_MN_SW64_Atom<ElementType>{}) == 0) {
|
||||
return UMMA::Layout_MN_SW64_Atom<ElementType>{};
|
||||
}
|
||||
else if constexpr (BLK_MN0 % size<0>(UMMA::Layout_MN_SW32_Atom<ElementType>{}) == 0) {
|
||||
return UMMA::Layout_MN_SW32_Atom<ElementType>{};
|
||||
}
|
||||
else if constexpr (BLK_MN0 % size<0>(UMMA::Layout_MN_INTER_Atom<ElementType>{}) == 0) {
|
||||
return UMMA::Layout_MN_INTER_Atom<ElementType>{};
|
||||
}
|
||||
else {
|
||||
static_assert(BLK_MN0 % size<0>(UMMA::Layout_MN_INTER_Atom<ElementType>{}) == 0,
|
||||
"BLK_MN0 must be a multiple of size<0>(UMMA::Layout_MN_INTER_Atom<ElementType>{})");
|
||||
}
|
||||
}
|
||||
}
|
||||
else if constexpr (major == cute::UMMA::Major::K) {
|
||||
if constexpr (BLK_K0 % size<1>(UMMA::Layout_K_SW128_Atom<ElementType>{}) == 0) {
|
||||
return UMMA::Layout_K_SW128_Atom<ElementType>{};
|
||||
}
|
||||
else if constexpr (BLK_K0 % size<1>(UMMA::Layout_K_SW64_Atom<ElementType>{}) == 0) {
|
||||
return UMMA::Layout_K_SW64_Atom<ElementType>{};
|
||||
}
|
||||
else if constexpr (BLK_K0 % size<1>(UMMA::Layout_K_SW32_Atom<ElementType>{}) == 0) {
|
||||
return UMMA::Layout_K_SW32_Atom<ElementType>{};
|
||||
}
|
||||
else if constexpr (BLK_K0 % size<1>(UMMA::Layout_K_INTER_Atom<ElementType>{}) == 0) {
|
||||
return UMMA::Layout_K_INTER_Atom<ElementType>{};
|
||||
}
|
||||
else {
|
||||
static_assert(BLK_K0 % size<1>(UMMA::Layout_K_INTER_Atom<ElementType>{}) == 0,
|
||||
"BLK_K0 must be a multiple of size<1>(UMMA::Layout_K_INTER_Atom<ElementType>{})");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
template <class ClusterShapeMNK, class AtomThrId>
|
||||
constexpr auto
|
||||
sm100_cluster_shape_to_tma_atom_A(ClusterShapeMNK cluster_shape_mnk, AtomThrId atom_thr_id) {
|
||||
static_assert(cute::rank(cluster_shape_mnk) == 3);
|
||||
constexpr bool IsDynamicCluster = not cute::is_static_v<ClusterShapeMNK>;
|
||||
|
||||
if constexpr (cute::size(atom_thr_id) == 2) {
|
||||
if constexpr (!IsDynamicCluster) {
|
||||
static_assert(cute::size<0>(cluster_shape_mnk) % 2 == 0, "Cluster shape not divisible by MMA size");
|
||||
if constexpr (cute::size<1>(cluster_shape_mnk) == 1) {
|
||||
return cute::SM100_TMA_2SM_LOAD{};
|
||||
}
|
||||
else {
|
||||
return cute::SM100_TMA_2SM_LOAD_MULTICAST{};
|
||||
}
|
||||
}
|
||||
else {
|
||||
return cute::SM100_TMA_2SM_LOAD_MULTICAST{};
|
||||
}
|
||||
}
|
||||
else if constexpr (size(atom_thr_id) == 1) {
|
||||
if constexpr (!IsDynamicCluster) {
|
||||
return detail::sm90_cluster_shape_to_tma_atom(cute::size<1>(cluster_shape_mnk));
|
||||
}
|
||||
else {
|
||||
// In the case of dynamic cluster, multicast decision is not known at compile time.
|
||||
// A multicast instruction is forced by passing a cute::Int<2>{} to this helper.
|
||||
return detail::sm90_cluster_shape_to_tma_atom(cute::Int<2>{});
|
||||
}
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<ClusterShapeMNK>,
|
||||
"Unsupported Configuration for SM100 TMA");
|
||||
}
|
||||
}
|
||||
|
||||
template <class ClusterShapeMNK, class AtomThrId>
|
||||
constexpr auto
|
||||
sm100_cluster_shape_to_tma_atom_B(ClusterShapeMNK cluster_shape_mnk, AtomThrId atom_thr_id) {
|
||||
static_assert(cute::rank(cluster_shape_mnk) == 3);
|
||||
constexpr bool IsDynamicCluster = not cute::is_static_v<ClusterShapeMNK>;
|
||||
|
||||
if constexpr (cute::size(atom_thr_id) == 2) {
|
||||
if constexpr (!IsDynamicCluster) {
|
||||
static_assert(cute::size<0>(cluster_shape_mnk) % 2 == 0, "Cluster shape not divisible by MMA size");
|
||||
if constexpr (cute::size<0>(cluster_shape_mnk) == 2) {
|
||||
return cute::SM100_TMA_2SM_LOAD{};
|
||||
}
|
||||
else {
|
||||
return cute::SM100_TMA_2SM_LOAD_MULTICAST{};
|
||||
}
|
||||
}
|
||||
else {
|
||||
return cute::SM100_TMA_2SM_LOAD_MULTICAST{};
|
||||
}
|
||||
} else if constexpr (size(atom_thr_id) == 1) {
|
||||
if constexpr (!IsDynamicCluster) {
|
||||
return detail::sm90_cluster_shape_to_tma_atom(cute::size<0>(cluster_shape_mnk));
|
||||
}
|
||||
else {
|
||||
// In the case of dynamic cluster, multicast decision is not known at compile time.
|
||||
// A multicast instruction is forced by passing a cute::Int<2>{} to this helper.
|
||||
return detail::sm90_cluster_shape_to_tma_atom(cute::Int<2>{});
|
||||
}
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<ClusterShapeMNK>,
|
||||
"Unsupported Configuration for SM100 TMA");
|
||||
}
|
||||
}
|
||||
|
||||
template<class KernelScheduleType>
|
||||
constexpr uint32_t find_vector_size() {
|
||||
if constexpr (cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecialized1SmNvf4Sm100> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecialized2SmNvf4Sm100> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelPtrArrayTmaWarpSpecialized1SmNvf4Sm100> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelPtrArrayTmaWarpSpecialized2SmNvf4Sm100>
|
||||
) {
|
||||
return 16;
|
||||
}
|
||||
else {
|
||||
return 32;
|
||||
}
|
||||
}
|
||||
|
||||
template<
|
||||
class ElementAMma,
|
||||
class ElementBMma,
|
||||
class ElementAMmaccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
UMMA::ScaleIn ANeg = UMMA::ScaleIn::One,
|
||||
UMMA::ScaleIn BNeg = UMMA::ScaleIn::One
|
||||
>
|
||||
constexpr auto
|
||||
sm100_make_1sm_trivial_tiled_mma() {
|
||||
|
||||
constexpr int M = cute::size<0>(TileShape_MNK{});
|
||||
static_assert(M == 64 || M == 128, "Invalid TileShape_M.");
|
||||
|
||||
// Do not allow a tiled MMA N mode > 1, as that is not reasonable.
|
||||
constexpr int N = cute::size<1>(TileShape_MNK{});
|
||||
static_assert(N % 8 == 0 && N <= 256, "Invalid TileShape_N.");
|
||||
|
||||
if constexpr (cute::is_same_v<ElementAMma, cutlass::tfloat32_t>) {
|
||||
static_assert(cute::is_same_v<ElementAMma, ElementBMma>, "ElementAMma and ElementBMma must match.");
|
||||
return make_tiled_mma(cute::SM100_MMA_TF32_SS<ElementAMma, ElementBMma, ElementAMmaccumulator,
|
||||
M, N, UmmaMajorA, UmmaMajorB, ANeg, BNeg>{});
|
||||
}
|
||||
else if constexpr (cute::is_same_v<ElementAMma, cutlass::half_t> ||
|
||||
cute::is_same_v<ElementAMma, cutlass::bfloat16_t>) {
|
||||
static_assert(cute::is_same_v<ElementAMma, ElementBMma>, "ElementAMma and ElementBMma must match.");
|
||||
return make_tiled_mma(cute::SM100_MMA_F16BF16_SS<ElementAMma, ElementBMma, ElementAMmaccumulator,
|
||||
M, N, UmmaMajorA, UmmaMajorB, ANeg, BNeg>{});
|
||||
}
|
||||
else if constexpr (cute::is_same_v<ElementAMma, int8_t> ||
|
||||
cute::is_same_v<ElementAMma, uint8_t>) {
|
||||
return make_tiled_mma(cute::SM100_MMA_S8_SS<ElementAMma, ElementBMma, ElementAMmaccumulator,
|
||||
M, N, UmmaMajorA, UmmaMajorB>{});
|
||||
}
|
||||
else if constexpr (cute::is_same_v<ElementAMma, cutlass::type_erased_dynamic_float8_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::detail::type_erased_dynamic_float6_unpacksmem_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::detail::type_erased_dynamic_float4_unpacksmem_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::float_e4m3_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::float_e5m2_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::detail::float_e2m3_unpacksmem_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::detail::float_e3m2_unpacksmem_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::detail::float_e2m1_unpacksmem_t>
|
||||
) {
|
||||
|
||||
return make_tiled_mma(
|
||||
cute::MMA_Traits<
|
||||
cute::SM100_MMA_F8F6F4_SS,
|
||||
ElementAMma,
|
||||
ElementBMma,
|
||||
ElementAMmaccumulator,
|
||||
cute::C<M>,
|
||||
cute::C<N>,
|
||||
cute::integral_constant<UMMA::Major, UmmaMajorA>,
|
||||
cute::integral_constant<UMMA::Major, UmmaMajorB>,
|
||||
cute::integral_constant<UMMA::ScaleIn, ANeg>,
|
||||
cute::integral_constant<UMMA::ScaleIn, BNeg>
|
||||
>{}
|
||||
);
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<ElementAMma>,
|
||||
"Unsupported configuration for SM100 collective builder.");
|
||||
}
|
||||
}
|
||||
|
||||
template<
|
||||
class ElementAMma,
|
||||
class ElementBMma,
|
||||
class ElementAMmaccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
UMMA::ScaleIn ANeg = UMMA::ScaleIn::One,
|
||||
UMMA::ScaleIn BNeg = UMMA::ScaleIn::One
|
||||
>
|
||||
constexpr auto
|
||||
sm100_make_2sm_trivial_tiled_mma() {
|
||||
|
||||
constexpr int M = cute::size<0>(TileShape_MNK{});
|
||||
static_assert(M == 128 || M == 256, "Invalid TileShape_M.");
|
||||
|
||||
// Do not allow a tiled MMA N mode > 1, as that is not reasonable.
|
||||
constexpr int N = cute::size<1>(TileShape_MNK{});
|
||||
static_assert(N % 8 == 0 && N <= 256, "Invalid TileShape_N.");
|
||||
|
||||
if constexpr (cute::is_same_v<ElementAMma, cutlass::tfloat32_t>) {
|
||||
static_assert(cute::is_same_v<ElementAMma, ElementBMma>, "ElementAMma and ElementBMma must match.");
|
||||
return make_tiled_mma(cute::SM100_MMA_TF32_2x1SM_SS<ElementAMma, ElementBMma, ElementAMmaccumulator,
|
||||
M, N, UmmaMajorA, UmmaMajorB, ANeg, BNeg>{});
|
||||
}
|
||||
else if constexpr (cute::is_same_v<ElementAMma, cutlass::half_t> ||
|
||||
cute::is_same_v<ElementAMma, cutlass::bfloat16_t>) {
|
||||
static_assert(cute::is_same_v<ElementAMma, ElementBMma>, "ElementAMma and ElementBMma must match.");
|
||||
return make_tiled_mma(cute::SM100_MMA_F16BF16_2x1SM_SS<ElementAMma, ElementBMma, ElementAMmaccumulator,
|
||||
M, N, UmmaMajorA, UmmaMajorB, ANeg, BNeg>{});
|
||||
}
|
||||
else if constexpr (cute::is_same_v<ElementAMma, int8_t> ||
|
||||
cute::is_same_v<ElementAMma, uint8_t>) {
|
||||
return make_tiled_mma(cute::SM100_MMA_S8_2x1SM_SS<ElementAMma, ElementBMma, ElementAMmaccumulator,
|
||||
M, N, UmmaMajorA, UmmaMajorB>{});
|
||||
}
|
||||
else if constexpr (cute::is_same_v<ElementAMma, cutlass::type_erased_dynamic_float8_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::detail::type_erased_dynamic_float6_unpacksmem_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::detail::type_erased_dynamic_float4_unpacksmem_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::float_e4m3_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::float_e5m2_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::detail::float_e2m3_unpacksmem_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::detail::float_e3m2_unpacksmem_t>
|
||||
|| cute::is_same_v<ElementAMma, cutlass::detail::float_e2m1_unpacksmem_t>
|
||||
) {
|
||||
|
||||
return make_tiled_mma(
|
||||
cute::MMA_Traits<
|
||||
cute::SM100_MMA_F8F6F4_2x1SM_SS,
|
||||
ElementAMma,
|
||||
ElementBMma,
|
||||
ElementAMmaccumulator,
|
||||
cute::C<M>,
|
||||
cute::C<N>,
|
||||
cute::integral_constant<UMMA::Major, UmmaMajorA>,
|
||||
cute::integral_constant<UMMA::Major, UmmaMajorB>,
|
||||
cute::integral_constant<UMMA::ScaleIn, ANeg>,
|
||||
cute::integral_constant<UMMA::ScaleIn, BNeg>
|
||||
>{}
|
||||
);
|
||||
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<ElementAMma>,
|
||||
"Unsupported configuration for SM100 collective builder.");
|
||||
}
|
||||
}
|
||||
|
||||
// For new MMA construction and partitioning that supports both dynamic and static cluster shape.
|
||||
// Used in conjunction with make_tma_atom_(A|B)_sm100
|
||||
// TileShape_MNK is always static and has shape (MmaAtomShapeM, MmaAtomShapeN, TileK)
|
||||
// ClusterShape_MNK can be dynamic or static.
|
||||
template<
|
||||
class ElementAMma,
|
||||
class ElementBMma,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
class KernelScheduleType,
|
||||
UMMA::ScaleIn ANeg = UMMA::ScaleIn::One,
|
||||
UMMA::ScaleIn BNeg = UMMA::ScaleIn::One
|
||||
>
|
||||
constexpr auto
|
||||
sm100_make_trivial_tiled_mma() {
|
||||
// MMA_2SM requested
|
||||
if constexpr (cute::is_base_of_v<KernelSchedule2Sm, KernelScheduleType> ) {
|
||||
return sm100_make_2sm_trivial_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
|
||||
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ANeg, BNeg>();
|
||||
}
|
||||
// MMA_1SM requested
|
||||
else if constexpr (cute::is_base_of_v<KernelSchedule1Sm, KernelScheduleType> ) {
|
||||
return sm100_make_1sm_trivial_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
|
||||
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ANeg, BNeg>();
|
||||
}
|
||||
// Auto scheduling requested
|
||||
else if constexpr (cute::is_same_v<KernelScheduleType, KernelScheduleAuto>) {
|
||||
// Static cluster
|
||||
if constexpr (cute::is_static_v<ClusterShape_MNK>) {
|
||||
// For MMA_2SM we need a cluster shape that is multiple of 2x1
|
||||
// and only M=128 and M=256 are supported, otherwise, fall back to MMA_1SM
|
||||
if constexpr (cute::size<0>(ClusterShape_MNK{}) % 2 == 0 &&
|
||||
cute::size<0>(TileShape_MNK{}) % 128 == 0) {
|
||||
return sm100_make_2sm_trivial_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
|
||||
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ANeg, BNeg>();
|
||||
}
|
||||
else {
|
||||
return sm100_make_1sm_trivial_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
|
||||
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ANeg, BNeg>();
|
||||
}
|
||||
// Dynamic cluster shape means we cannot assume we can use 2SM MMA
|
||||
}
|
||||
else {
|
||||
return sm100_make_1sm_trivial_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator,
|
||||
TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ANeg, BNeg>();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
/**
|
||||
* @brief Check for U4_UNPACK_U8, U6_UNPACK_U8 alignment requirement
|
||||
*
|
||||
* @tparam TileShape_MNK (MmaAtomShape_M, MmaAtomShape_N, TileShape_K)
|
||||
* @tparam ClusterShape_MNK (cluster_M, cluster_N, cluster_K)
|
||||
* @tparam KernelScheduleType Builder tag
|
||||
*/
|
||||
template<
|
||||
class ElementAMma,
|
||||
class ElementBMma,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
class KernelScheduleType,
|
||||
bool Is2sm
|
||||
>
|
||||
constexpr bool sm100_gemm_check_for_f8f6f4_mix8bit_requirement(){
|
||||
|
||||
|
||||
[[maybe_unused]] constexpr int TileShape_M = Is2sm ? size<0>(TileShape_MNK{}) / 2 : size<0>(TileShape_MNK{});
|
||||
[[maybe_unused]] constexpr int TileShape_N = size<1>(TileShape_MNK{});
|
||||
[[maybe_unused]] constexpr int TileShape_K = size<2>(TileShape_MNK{});
|
||||
|
||||
constexpr bool is_b_unpack_f4_f6 = cute::is_same_v<ElementBMma, cutlass::detail::float_e2m1_unpacksmem_t> ||
|
||||
cute::is_same_v<ElementBMma, cutlass::detail::float_e3m2_unpacksmem_t> ||
|
||||
cute::is_same_v<ElementBMma, cutlass::detail::float_e2m3_unpacksmem_t> ||
|
||||
cute::is_same_v<ElementBMma, cutlass::detail::type_erased_dynamic_float4_unpacksmem_t> ||
|
||||
cute::is_same_v<ElementBMma, cutlass::detail::type_erased_dynamic_float6_unpacksmem_t>;
|
||||
constexpr bool is_a_unpack_f4_f6 = cute::is_same_v<ElementAMma, cutlass::detail::float_e2m1_unpacksmem_t> ||
|
||||
cute::is_same_v<ElementAMma, cutlass::detail::float_e3m2_unpacksmem_t> ||
|
||||
cute::is_same_v<ElementAMma, cutlass::detail::float_e2m3_unpacksmem_t> ||
|
||||
cute::is_same_v<ElementAMma, cutlass::detail::type_erased_dynamic_float4_unpacksmem_t> ||
|
||||
cute::is_same_v<ElementAMma, cutlass::detail::type_erased_dynamic_float6_unpacksmem_t>;
|
||||
|
||||
[[maybe_unused]] constexpr bool is_b_n_major = UmmaMajorB == UMMA::Major::MN;
|
||||
[[maybe_unused]] constexpr bool is_b_k_major = !is_b_n_major;
|
||||
[[maybe_unused]] constexpr bool is_a_m_major = UmmaMajorA == UMMA::Major::MN;
|
||||
[[maybe_unused]] constexpr bool is_a_k_major = !is_a_m_major;
|
||||
|
||||
// 2SM
|
||||
if constexpr (Is2sm) {
|
||||
constexpr bool valid_a = !is_a_unpack_f4_f6 || (is_a_k_major ?
|
||||
TileShape_K % 128 == 0 :
|
||||
TileShape_M % 128 == 0);
|
||||
|
||||
constexpr bool valid_b = !is_b_unpack_f4_f6 || (is_b_n_major ?
|
||||
TileShape_N % 256 == 0:
|
||||
TileShape_K % 128 == 0);
|
||||
return valid_a && valid_b;
|
||||
}
|
||||
// 1SM
|
||||
else {
|
||||
constexpr bool valid_a = !is_a_unpack_f4_f6 || (is_a_k_major ?
|
||||
TileShape_K % 128 == 0 :
|
||||
TileShape_M % 128 == 0);
|
||||
|
||||
constexpr bool valid_b = !is_b_unpack_f4_f6 || (is_b_n_major ?
|
||||
TileShape_N % 128 == 0 :
|
||||
TileShape_K % 128 == 0);
|
||||
|
||||
return valid_a && valid_b;
|
||||
}
|
||||
}
|
||||
|
||||
template <class ElementA, int AlignmentA, class ElementB, int AlignmentB, class KernelScheduleType>
|
||||
constexpr bool
|
||||
sm1xx_gemm_is_aligned() {
|
||||
// Only support dense gemm alignment check
|
||||
constexpr bool is_f6f4_subbytes = cute::sizeof_bits_v<ElementA> < 8 || cute::sizeof_bits_v<ElementB> < 8;
|
||||
|
||||
return ((cute::sizeof_bits_v<ElementA> * AlignmentA) % cutlass::detail::get_input_alignment_bits<ElementA, is_f6f4_subbytes>() == 0) &&
|
||||
((cute::sizeof_bits_v<ElementB> * AlignmentB) % cutlass::detail::get_input_alignment_bits<ElementB, is_f6f4_subbytes>() == 0);
|
||||
}
|
||||
|
||||
template <class ElementA, int AlignmentA, class ElementB, int AlignmentB, class KernelScheduleType>
|
||||
constexpr bool
|
||||
sm1xx_blockscaled_gemm_is_aligned() {
|
||||
// Only support blocksscaled gemm alignment check
|
||||
constexpr bool is_f6f4_subbytes = (cute::sizeof_bits_v<ElementA> < 8 || cute::sizeof_bits_v<ElementB> < 8) &&
|
||||
(cute::is_base_of_v<KernelScheduleMxf8f6f4Sm100, KernelScheduleType>
|
||||
);
|
||||
|
||||
return ((cute::sizeof_bits_v<ElementA> * AlignmentA) % cutlass::detail::get_input_alignment_bits<ElementA, is_f6f4_subbytes>() == 0) &&
|
||||
((cute::sizeof_bits_v<ElementB> * AlignmentB) % cutlass::detail::get_input_alignment_bits<ElementB, is_f6f4_subbytes>() == 0);
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::collective
|
||||
@@ -0,0 +1,117 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
|
||||
|
||||
#pragma once
|
||||
|
||||
namespace cutlass::gemm::collective::detail {
|
||||
|
||||
template<
|
||||
class ClusterShape_MNK,
|
||||
int AccumulatorPipelineStageCount,
|
||||
int SchedulerPipelineStageCount,
|
||||
int CLCResponseSize,
|
||||
bool IsArrayOfPointersGemm,
|
||||
int NumTensorMaps=2
|
||||
>
|
||||
struct Sm100DenseGemmTmaUmmaCarveout {
|
||||
// AccumulatorPipeline = PipelineUmmaAsync
|
||||
static constexpr auto AccumulatorPipelineStorage = sizeof(typename cutlass::PipelineUmmaAsync<AccumulatorPipelineStageCount>::SharedStorage);
|
||||
// CLCPipeline = PipelineCLCFetchAsync
|
||||
static constexpr auto CLCPipelineStorage = sizeof(typename cutlass::PipelineCLCFetchAsync<SchedulerPipelineStageCount, ClusterShape_MNK>::SharedStorage);
|
||||
// LoadOrderBarrier = OrderedSequenceBarrier<1,2>
|
||||
static constexpr auto LoadOrderBarrierStorage = sizeof(typename cutlass::OrderedSequenceBarrier<1,2>::SharedStorage);
|
||||
// CLC (scheduler) response
|
||||
static constexpr auto CLCResponseStorage = SchedulerPipelineStageCount * detail::CLCResponseSize;
|
||||
// CLC Throttle pipeline storage
|
||||
static constexpr auto CLCThrottlePipelineStorage = sizeof(typename cutlass::PipelineAsync<SchedulerPipelineStageCount>::SharedStorage);
|
||||
// Tmem dealloc
|
||||
static constexpr auto TmemDeallocStorage = sizeof(cutlass::arch::ClusterBarrier);
|
||||
// Tmem ptr storage
|
||||
static constexpr auto TmemBasePtrsStorage = SchedulerPipelineStageCount * sizeof(uint32_t);
|
||||
// Tensormap Storage
|
||||
static constexpr auto TensorMapStorage =
|
||||
IsArrayOfPointersGemm ? sizeof(cute::TmaDescriptor) * NumTensorMaps /* for A and B */ :
|
||||
0;
|
||||
// Smem usage that's not part of CollectiveEpilogue::SharedStorage & CollectiveMainloop::SharedStorage
|
||||
static constexpr auto KernelSmemCarveout = static_cast<int>( AccumulatorPipelineStorage +
|
||||
CLCPipelineStorage +
|
||||
LoadOrderBarrierStorage +
|
||||
TmemDeallocStorage +
|
||||
CLCThrottlePipelineStorage +
|
||||
CLCResponseStorage +
|
||||
TmemBasePtrsStorage +
|
||||
TensorMapStorage
|
||||
);
|
||||
};
|
||||
|
||||
template<class ClusterShape_MNK, int AccumulatorPipelineStageCount, int SchedulerPipelineStageCount, int CLCResponseSize>
|
||||
struct Sm100SparseGemmTmaUmmaCarveout {
|
||||
|
||||
// * GemmUniversal::SharedStorage::PipelineStorage
|
||||
// LoadOrderBarrier = OrderedSequenceBarrier<1,2>
|
||||
static constexpr auto LoadOrderBarrierStorage = sizeof(typename cutlass::OrderedSequenceBarrier<1,2>::SharedStorage);
|
||||
// CLCPipelineStorage = PipelineCLCFetchAsync
|
||||
static constexpr auto CLCPipelineStorage = sizeof(typename cutlass::PipelineCLCFetchAsync<SchedulerPipelineStageCount, ClusterShape_MNK>::SharedStorage);
|
||||
// AccumulatorPipeline = PipelineUmmaAsync
|
||||
static constexpr auto AccumulatorPipelineStorage = sizeof(typename cutlass::PipelineUmmaAsync<AccumulatorPipelineStageCount>::SharedStorage);
|
||||
// CLC Throttle pipeline storage
|
||||
static constexpr auto CLCThrottlePipelineStorage = sizeof(typename cutlass::PipelineAsync<SchedulerPipelineStageCount>::SharedStorage);
|
||||
// Tmem dealloc
|
||||
static constexpr auto TmemDeallocStorage = sizeof(cutlass::arch::ClusterBarrier);
|
||||
// Epilogue Throttle
|
||||
static constexpr auto EpilogueThrottleStorage = sizeof(arch::ClusterBarrier);
|
||||
|
||||
static constexpr auto PipelineStorage = static_cast<int>(cutlass::round_up(
|
||||
cutlass::round_up(LoadOrderBarrierStorage, 16) +
|
||||
cutlass::round_up(CLCPipelineStorage, 16) +
|
||||
cutlass::round_up(AccumulatorPipelineStorage, 16) +
|
||||
cutlass::round_up(CLCThrottlePipelineStorage, 16) +
|
||||
cutlass::round_up(TmemDeallocStorage, 8) +
|
||||
cutlass::round_up(EpilogueThrottleStorage, 8),
|
||||
16));
|
||||
|
||||
// * GemmUniversal::SharedStorage::Others
|
||||
// CLC (scheduler) response
|
||||
static constexpr auto CLCQueryResponseStorage = SchedulerPipelineStageCount * CLCResponseSize;
|
||||
// Tmem ptr storage
|
||||
static constexpr auto TmemBasePtrsStorage = sizeof(uint32_t);
|
||||
|
||||
static constexpr auto OtherStorage = static_cast<int>(cutlass::round_up(
|
||||
cutlass::round_up(CLCQueryResponseStorage, 16) +
|
||||
cutlass::round_up(TmemBasePtrsStorage, 16),
|
||||
16));
|
||||
|
||||
// Smem usage that's not part of CollectiveEpilogue::SharedStorage & CollectiveMainloop::SharedStorage
|
||||
static constexpr auto KernelSmemCarveout = static_cast<int>( PipelineStorage +
|
||||
OtherStorage);
|
||||
};
|
||||
} // namespace cutlass::gemm::collective::detail
|
||||
@@ -0,0 +1,320 @@
|
||||
/***************************************************************************************************
|
||||
* 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"
|
||||
#include "cutlass/gemm/collective/builders/sm100_pipeline_carveout.inl"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::collective {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
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,
|
||||
class MainloopPipelineStorage,
|
||||
int stages
|
||||
>
|
||||
constexpr int
|
||||
sm100_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,
|
||||
class MainloopPipelineStorage,
|
||||
int stages
|
||||
>
|
||||
constexpr int
|
||||
sm100_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,
|
||||
class MainloopPipelineStorage,
|
||||
int carveout_bytes>
|
||||
constexpr int
|
||||
sm100_compute_stage_count_or_override(StageCountAutoCarveout<carveout_bytes> stage_count) {
|
||||
// For F8F6F4 sub-bytes, ElementA/B will be passed in as uint8_t
|
||||
// For Planar Complex, ElementA/B will be passed in as cutlass::complex<ElementARaw>
|
||||
// Each stage include (CollectiveMma::SharedStorage)
|
||||
// 1. smem for A and smem for B (CollectiveMma::SharedStorage::TensorStorage)
|
||||
// 2. one MainloopPipeline = (CollectiveMma::SharedStorage::PipelineStorage = PipelineTmaUmmaAsync)
|
||||
constexpr auto mainloop_pipeline_bytes = sizeof(MainloopPipelineStorage);
|
||||
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;
|
||||
}
|
||||
|
||||
template <class ElementA, class ElementB>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static constexpr bool
|
||||
check_input_datatypes() {
|
||||
auto is_non_f4f6f8_input = [&]() {
|
||||
return (cute::is_same_v<ElementA, cutlass::tfloat32_t> ||
|
||||
cute::is_same_v<ElementA, float> ||
|
||||
cute::is_same_v<ElementA, cutlass::half_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::bfloat16_t> ||
|
||||
cute::is_same_v<ElementA, int8_t> ||
|
||||
cute::is_same_v<ElementA, uint8_t>) &&
|
||||
(cute::is_same_v<ElementA, ElementB>); // For all MMA instrs except F4F6F8, A and B types should be the same.
|
||||
};
|
||||
auto is_f4f6f8_input = [&]() {
|
||||
// Allowed input element datatype for narrow precision GEMM
|
||||
return (
|
||||
(
|
||||
cute::is_same_v<ElementA, cutlass::type_erased_dynamic_float8_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::type_erased_dynamic_float6_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::type_erased_dynamic_float4_t>
|
||||
) &&
|
||||
(
|
||||
cute::is_same_v<ElementB, cutlass::type_erased_dynamic_float8_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::type_erased_dynamic_float6_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::type_erased_dynamic_float4_t>
|
||||
)
|
||||
) ||
|
||||
(
|
||||
(
|
||||
cute::is_same_v<ElementA, cutlass::float_e2m1_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::float_e2m3_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::float_e3m2_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::float_e4m3_t> ||
|
||||
cute::is_same_v<ElementA, cutlass::float_e5m2_t>
|
||||
) &&
|
||||
(
|
||||
cute::is_same_v<ElementB, cutlass::float_e2m1_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::float_e2m3_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::float_e3m2_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::float_e4m3_t> ||
|
||||
cute::is_same_v<ElementB, cutlass::float_e5m2_t>
|
||||
)
|
||||
);
|
||||
};
|
||||
|
||||
static_assert(is_f4f6f8_input() || is_non_f4f6f8_input(), "Unsupported data type for ElementA");
|
||||
|
||||
return true;
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
template <
|
||||
class ElementA,
|
||||
class GmemLayoutATag,
|
||||
int AlignmentA,
|
||||
class ElementB,
|
||||
class GmemLayoutBTag,
|
||||
int AlignmentB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
class StageCountType,
|
||||
class KernelScheduleType
|
||||
>
|
||||
struct CollectiveBuilder<
|
||||
arch::Sm100,
|
||||
arch::OpClassTensorOp,
|
||||
ElementA,
|
||||
GmemLayoutATag,
|
||||
AlignmentA,
|
||||
ElementB,
|
||||
GmemLayoutBTag,
|
||||
AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape_MNK, // (MmaAtomShapeM, MmaAtomShapeN, TileK)
|
||||
ClusterShape_MNK, // Static cluster shape or dynamic (int, int, _1)
|
||||
StageCountType,
|
||||
KernelScheduleType,
|
||||
cute::enable_if_t<
|
||||
not cute::is_tuple_v<ElementA> && not cute::is_tuple_v<ElementB> &&
|
||||
not cute::is_complex_v<ElementA> && not cute::is_complex_v<ElementB> &&
|
||||
// Dense Gemm / PtrArrayDenseGemm
|
||||
(
|
||||
(cute::is_base_of_v<KernelScheduleSm100DenseGemm, KernelScheduleType> ||
|
||||
cute::is_same_v<KernelScheduleAuto, KernelScheduleType>)) &&
|
||||
// Alignment check
|
||||
detail::sm1xx_gemm_is_aligned<ElementA, AlignmentA, ElementB, AlignmentB, KernelScheduleType>()>>
|
||||
{
|
||||
static_assert(cute::is_static_v<TileShape_MNK>, "TileShape has to be static");
|
||||
static_assert(detail::check_input_datatypes<ElementA, ElementB>(), "Incorrect input types");
|
||||
|
||||
static constexpr cute::UMMA::Major UmmaMajorA = cutlass::gemm::collective::detail::tag_to_umma_major_A<GmemLayoutATag>();
|
||||
static constexpr cute::UMMA::Major UmmaMajorB = cutlass::gemm::collective::detail::tag_to_umma_major_B<GmemLayoutBTag>();
|
||||
|
||||
// Data type used by MMA instruction
|
||||
using ElementAMma = decltype(cutlass::gemm::collective::detail::sm100_kernel_input_element_to_mma_input_element<ElementA>());
|
||||
using ElementBMma = decltype(cutlass::gemm::collective::detail::sm100_kernel_input_element_to_mma_input_element<ElementB>());
|
||||
|
||||
static constexpr bool is_2sm = cute::is_base_of_v<KernelSchedule2Sm, KernelScheduleType> ||
|
||||
(not cute::is_base_of_v<KernelSchedule1Sm, KernelScheduleType> &&
|
||||
not cute::is_base_of_v<KernelSchedule2Sm, KernelScheduleType> &&
|
||||
cute::is_static_v<ClusterShape_MNK> &&
|
||||
cute::get<0>(ClusterShape_MNK{}) % 2 == 0 );
|
||||
|
||||
static_assert(detail::sm100_gemm_check_for_f8f6f4_mix8bit_requirement<ElementAMma, ElementBMma,
|
||||
TileShape_MNK, ClusterShape_MNK,
|
||||
UmmaMajorA, UmmaMajorB, KernelScheduleType, is_2sm>(),
|
||||
"TileSize and MNK Major does not met with MMA Mix 8-bit TMA load requirement" );
|
||||
using TiledMma = decltype(detail::sm100_make_trivial_tiled_mma<
|
||||
ElementAMma, ElementBMma, ElementAccumulator,
|
||||
decltype(cute::product_each(TileShape_MNK{})), ClusterShape_MNK,
|
||||
UmmaMajorA, UmmaMajorB, KernelScheduleType>());
|
||||
|
||||
using ElementAMma_SmemAllocType = cute::conditional_t<cute::sizeof_bits_v<ElementAMma> < 8, uint8_t, ElementAMma>;
|
||||
using ElementBMma_SmemAllocType = cute::conditional_t<cute::sizeof_bits_v<ElementBMma> < 8, uint8_t, ElementBMma>;
|
||||
|
||||
using AtomThrID = typename TiledMma::AtomThrID;
|
||||
|
||||
using AtomThrShapeMNK = cute::Shape<decltype(cute::shape<0>(typename TiledMma::ThrLayoutVMNK{})), _1, _1>;
|
||||
using CtaTileShape_MNK = decltype(cute::shape_div(TileShape_MNK{}, AtomThrShapeMNK{}));
|
||||
|
||||
// ((MMA_TILE_M,MMA_TILE_K), MMA_M, MMA_K)
|
||||
using MmaShapeA_MK = decltype(partition_shape_A(TiledMma{}, make_shape(cute::size<0>(TileShape_MNK{}),
|
||||
cute::size<2>(TileShape_MNK{}))));
|
||||
// ((MMA_TILE_N,MMA_TILE_K), MMA_N, MMA_K)
|
||||
using MmaShapeB_NK = decltype(partition_shape_B(TiledMma{}, make_shape(cute::size<1>(TileShape_MNK{}),
|
||||
cute::size<2>(TileShape_MNK{}))));
|
||||
|
||||
using BlockTileA_M = decltype(cute::size<0,0>(MmaShapeA_MK{}) * cute::size<1>(MmaShapeA_MK{}));
|
||||
using BlockTileA_K = decltype(cute::size<0,1>(MmaShapeA_MK{}) * cute::size<2>(MmaShapeA_MK{}));
|
||||
using BlockTileB_N = decltype(cute::size<0,0>(MmaShapeB_NK{}) * cute::size<1>(MmaShapeB_NK{}));
|
||||
using BlockTileB_K = decltype(cute::size<0,1>(MmaShapeB_NK{}) * cute::size<2>(MmaShapeB_NK{}));
|
||||
|
||||
// Kludged right divide to divide TileShape_M/N by 1SM/2SM
|
||||
// Future work: fix partition_shape to account for hierarchies and
|
||||
// contiguity so we can pass BlockTileA/B to sm100_smem_selector instead
|
||||
using SmemShape_M = decltype(shape_div(shape<0>(TileShape_MNK{}), shape_div(shape<0>(TileShape_MNK{}), size<0>(TileShape_MNK{}) / size(AtomThrID{}))));
|
||||
using SmemShape_N = decltype(shape_div(shape<1>(TileShape_MNK{}), shape_div(shape<1>(TileShape_MNK{}), size<1>(TileShape_MNK{}) / size(AtomThrID{}))));
|
||||
using SmemShape_K = decltype(cute::get<2>(TileShape_MNK{}));
|
||||
|
||||
using GmemTiledCopyA = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_A(
|
||||
ClusterShape_MNK{}, AtomThrID{}));
|
||||
using GmemTiledCopyB = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_B(
|
||||
ClusterShape_MNK{}, AtomThrID{}));
|
||||
|
||||
using SmemLayoutAtomA = decltype(cutlass::gemm::collective::detail::sm100_smem_selector<
|
||||
UmmaMajorA, ElementAMma_SmemAllocType, SmemShape_M, SmemShape_K>());
|
||||
using SmemLayoutAtomB = decltype(cutlass::gemm::collective::detail::sm100_smem_selector<
|
||||
UmmaMajorB, ElementBMma_SmemAllocType, SmemShape_N, SmemShape_K>());
|
||||
static constexpr uint32_t TotalTmemRows = 128;
|
||||
static constexpr uint32_t Sm100TmemCapacityColumns = 512;
|
||||
static constexpr uint32_t TotalTmem = TotalTmemRows * Sm100TmemCapacityColumns;
|
||||
static constexpr uint32_t AccumulatorPipelineStageCount = TotalTmem / (cute::size<0>(CtaTileShape_MNK{}) * cute::size<1>(CtaTileShape_MNK{}));
|
||||
static_assert(AccumulatorPipelineStageCount > 0, "Accumulator pipeline stage count must be positive. This error probably means that TileShape_MNK and/or TiledMma::ThrLayoutVMNK are wrong.");
|
||||
|
||||
// Calculate scheduler pipeline stages. Having one more stage than the accumulator allows more latency hiding.
|
||||
using StrideA = cutlass::gemm::TagToStrideA_t<GmemLayoutATag>;
|
||||
using InternalStrideA = cute::remove_pointer_t<StrideA>;
|
||||
// Grouped GEMM (where Stride type is Stride*) does not use CLC based scheduler.
|
||||
// SchedulerPipelineStageCount could be set to zero for Grouped GEMM, but we shouldn't define CLC Pipeline's barrier arrays of size zero.
|
||||
static constexpr uint32_t SchedulerPipelineStageCount = cute::is_same_v<InternalStrideA, StrideA> ? (AccumulatorPipelineStageCount + 1) : 1;
|
||||
static constexpr bool IsArrayOfPointersGemm = (cute::is_base_of_v<KernelScheduleSm100PtrArrayDenseGemm, KernelScheduleType>);
|
||||
static constexpr uint32_t KernelSmemCarveout = detail::Sm100DenseGemmTmaUmmaCarveout<
|
||||
ClusterShape_MNK,
|
||||
AccumulatorPipelineStageCount,
|
||||
SchedulerPipelineStageCount,
|
||||
detail::CLCResponseSize,
|
||||
IsArrayOfPointersGemm
|
||||
>::KernelSmemCarveout;
|
||||
// Reduce SMEM capacity available for buffers considering barrier allocations.
|
||||
static constexpr int Sm100ReducedSmemCapacityBytes = cutlass::gemm::collective::detail::sm100_smem_capacity_bytes - KernelSmemCarveout;
|
||||
|
||||
using SmemTileShape = cute::Shape<BlockTileA_M, BlockTileB_N, BlockTileA_K>;
|
||||
using MainloopPipelineStorage = typename cutlass::PipelineTmaUmmaAsync<1>::SharedStorage;
|
||||
|
||||
static constexpr int PipelineStages = cutlass::gemm::collective::detail::sm100_compute_stage_count_or_override<
|
||||
Sm100ReducedSmemCapacityBytes, ElementAMma_SmemAllocType, ElementBMma_SmemAllocType, SmemTileShape, MainloopPipelineStorage>(StageCountType{});
|
||||
static_assert(PipelineStages > 0, "Smem usage is too high. Can't create any SMEM buffers for A, and B.");
|
||||
|
||||
using DispatchPolicy =
|
||||
cute::conditional_t<IsArrayOfPointersGemm,
|
||||
cutlass::gemm::MainloopSm100ArrayTmaUmmaWarpSpecialized<
|
||||
PipelineStages,
|
||||
SchedulerPipelineStageCount,
|
||||
AccumulatorPipelineStageCount,
|
||||
ClusterShape_MNK
|
||||
>,
|
||||
cutlass::gemm::MainloopSm100TmaUmmaWarpSpecialized<
|
||||
PipelineStages,
|
||||
SchedulerPipelineStageCount,
|
||||
AccumulatorPipelineStageCount,
|
||||
ClusterShape_MNK
|
||||
>
|
||||
>;
|
||||
|
||||
using CollectiveOp = cutlass::gemm::collective::CollectiveMma<
|
||||
DispatchPolicy,
|
||||
TileShape_MNK,
|
||||
ElementA,
|
||||
cutlass::gemm::TagToStrideA_t<GmemLayoutATag>,
|
||||
ElementB,
|
||||
cutlass::gemm::TagToStrideB_t<GmemLayoutBTag>,
|
||||
TiledMma,
|
||||
GmemTiledCopyA,
|
||||
SmemLayoutAtomA,
|
||||
void,
|
||||
cute::identity,
|
||||
GmemTiledCopyB,
|
||||
SmemLayoutAtomB,
|
||||
void,
|
||||
cute::identity
|
||||
>;
|
||||
};
|
||||
|
||||
} // namespace cutlass::gemm::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -39,4 +39,9 @@
|
||||
#include "cutlass/gemm/collective/collective_builder_decl.hpp"
|
||||
#include "cutlass/gemm/collective/builders/sm90_gmma_builder.inl"
|
||||
#include "cutlass/gemm/collective/builders/sm90_sparse_gmma_builder.inl"
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
#include "cutlass/gemm/collective/builders/sm100_umma_builder.inl"
|
||||
#include "cutlass/gemm/collective/builders/sm100_blockscaled_umma_builder.inl"
|
||||
#endif
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -48,4 +48,10 @@
|
||||
#include "cutlass/gemm/collective/sm90_mma_array_tma_gmma_ss_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_tma_gmma_ss_warpspecialized_fp8.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_tma_gmma_ss_warpspecialized_fp8_blockwise_scaling.hpp"
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
#include "cutlass/gemm/collective/sm100_mma_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm100_mma_array_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm100_blockscaled_mma_array_warpspecialized.hpp"
|
||||
#endif // !defined(__CUDACC_RTC__)
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,864 @@
|
||||
/***************************************************************************************************
|
||||
* 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/detail/collective.hpp"
|
||||
#include "cutlass/detail/cluster.hpp"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/trace.h"
|
||||
#include "cutlass/kernel_hardware_info.hpp"
|
||||
#include "cutlass/cuda_host_adapter.hpp"
|
||||
|
||||
#include "cute/algorithm/functional.hpp"
|
||||
#include "cute/arch/cluster_sm90.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cute/algorithm/gemm.hpp"
|
||||
#include "cute/tensor_predicate.hpp"
|
||||
#include "cute/numeric/arithmetic_tuple.hpp"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::collective {
|
||||
using namespace cute;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// WarpSpecialized Mainloop
|
||||
// Both DMA Load and MMA methods of this class must be run by a single thread that's picked by elect_one
|
||||
template <
|
||||
int Stages,
|
||||
int SchedulerPipelineStageCount,
|
||||
int AccumulatorPipelineStageCount,
|
||||
class ClusterShape, // Static cluster shape or dynamic (int, int, _1)
|
||||
class TileShape_, // (MmaAtomShapeM, MmaAtomShapeN, TileK)
|
||||
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<
|
||||
MainloopSm100ArrayTmaUmmaWarpSpecialized<
|
||||
Stages,
|
||||
SchedulerPipelineStageCount,
|
||||
AccumulatorPipelineStageCount,
|
||||
ClusterShape>,
|
||||
TileShape_,
|
||||
ElementA_,
|
||||
StrideA_,
|
||||
ElementB_,
|
||||
StrideB_,
|
||||
TiledMma_,
|
||||
GmemTiledCopyA_,
|
||||
SmemLayoutAtomA_,
|
||||
SmemCopyAtomA_,
|
||||
TransformA_,
|
||||
GmemTiledCopyB_,
|
||||
SmemLayoutAtomB_,
|
||||
SmemCopyAtomB_,
|
||||
TransformB_>
|
||||
{
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
using TiledMma = TiledMma_;
|
||||
using AtomThrShapeMNK = Shape<decltype(shape<0>(typename TiledMma::ThrLayoutVMNK{})), _1, _1>;
|
||||
|
||||
using DispatchPolicy = MainloopSm100ArrayTmaUmmaWarpSpecialized<
|
||||
Stages,
|
||||
SchedulerPipelineStageCount,
|
||||
AccumulatorPipelineStageCount,
|
||||
ClusterShape>;
|
||||
using TileShape = TileShape_;
|
||||
|
||||
static constexpr bool IsDynamicCluster = not cute::is_static_v<ClusterShape>;
|
||||
|
||||
CUTE_STATIC_ASSERT_V(evenly_divides(TileShape{}, tile_shape(TiledMma{})),
|
||||
"Static cluster shape used: TileShape should be evenly divided by TiledMma");
|
||||
|
||||
using CtaShape_MNK = decltype(shape_div(TileShape{}, AtomThrShapeMNK{}));
|
||||
|
||||
// Define A and B block shapes for reduced size TMA_LOADs
|
||||
using MmaShapeA_MK = decltype(partition_shape_A(TiledMma{}, make_shape(size<0>(TileShape{}), size<2>(TileShape{}))));
|
||||
using MmaShapeB_NK = decltype(partition_shape_B(TiledMma{}, make_shape(size<1>(TileShape{}), size<2>(TileShape{}))));
|
||||
|
||||
using ElementA = ElementA_;
|
||||
using ElementAMma = typename TiledMma::ValTypeA;
|
||||
using StrideA = StrideA_;
|
||||
using InternalStrideA = cute::remove_pointer_t<StrideA>;
|
||||
using ElementB = ElementB_;
|
||||
using ElementBMma = typename TiledMma::ValTypeB;
|
||||
using StrideB = StrideB_;
|
||||
using InternalStrideB = cute::remove_pointer_t<StrideB>;
|
||||
|
||||
static constexpr bool IsRuntimeDataTypeA = cutlass::gemm::collective::detail::is_sm10x_runtime_f8f6f4<ElementA>();
|
||||
|
||||
static constexpr bool IsRuntimeDataTypeB = cutlass::gemm::collective::detail::is_sm10x_runtime_f8f6f4<ElementB>();
|
||||
|
||||
static_assert((IsRuntimeDataTypeA && IsRuntimeDataTypeB) ||
|
||||
(!IsRuntimeDataTypeA && !IsRuntimeDataTypeB),
|
||||
"ElementA and ElementB should be both runtime or both static.");
|
||||
|
||||
static constexpr bool IsRuntimeDataType = IsRuntimeDataTypeA && IsRuntimeDataTypeB;
|
||||
|
||||
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 MainloopPipeline = cutlass::PipelineTmaUmmaAsync<
|
||||
DispatchPolicy::Stages,
|
||||
ClusterShape,
|
||||
AtomThrShapeMNK>;
|
||||
using MainloopPipelineState = typename MainloopPipeline::PipelineState;
|
||||
|
||||
static_assert(rank(SmemLayoutAtomA{}) == 2, "SmemLayoutAtomA must be rank 2 (M,K)");
|
||||
static_assert(((size<0,0>(MmaShapeA_MK{}) * size<1>(MmaShapeA_MK{})) % size<0>(SmemLayoutAtomA{})) == 0,
|
||||
"SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert(((size<0,1>(MmaShapeA_MK{}) * size<2>(MmaShapeA_MK{})) % size<1>(SmemLayoutAtomA{})) == 0,
|
||||
"SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert(cute::is_void_v<SmemCopyAtomA>,
|
||||
"SM100 UMMA cannot have a non-void copy atom for smem sourced instructions.");
|
||||
|
||||
static_assert(rank(SmemLayoutAtomB{}) == 2, "SmemLayoutAtomB must be rank 2 (N,K)");
|
||||
static_assert(((size<0,0>(MmaShapeB_NK{}) * size<1>(MmaShapeB_NK{})) % size<0>(SmemLayoutAtomB{})) == 0,
|
||||
"SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert(((size<0,1>(MmaShapeB_NK{}) * size<2>(MmaShapeB_NK{})) % size<1>(SmemLayoutAtomB{})) == 0,
|
||||
"SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert(cute::is_void_v<SmemCopyAtomB>,
|
||||
"SM100 UMMA cannot have a non-void copy atom for smem sourced instructions.");
|
||||
|
||||
// Tile along K mode first before tiling over MN. PIPE mode last as usual.
|
||||
// This maximizes TMA boxes due to better smem-K vectorization, reducing total issued TMAs.
|
||||
// (MMA_TILE_M,MMA_TILE_K),MMA_M,MMA_K,PIPE)
|
||||
using SmemLayoutA = decltype(UMMA::tile_to_mma_shape(
|
||||
SmemLayoutAtomA{},
|
||||
append(MmaShapeA_MK{}, Int<DispatchPolicy::Stages>{}),
|
||||
cute::conditional_t<cutlass::gemm::detail::is_mn_major<InternalStrideA>(), Step<_2,_1,_3>, Step<_1,_2,_3>>{}));
|
||||
// (MMA_TILE_N,MMA_TILE_K),MMA_N,MMA_K,PIPE)
|
||||
using SmemLayoutB = decltype(UMMA::tile_to_mma_shape(
|
||||
SmemLayoutAtomB{},
|
||||
append(MmaShapeB_NK{}, Int<DispatchPolicy::Stages>{}),
|
||||
cute::conditional_t<cutlass::gemm::detail::is_mn_major<InternalStrideB>(), Step<_2,_1,_3>, Step<_1,_2,_3>>{}));
|
||||
|
||||
static_assert(DispatchPolicy::Stages >= 2, "Specialization requires Stages set to value 1 or more.");
|
||||
static_assert(cute::is_base_of<cute::UMMA::DescriptorIterator, typename TiledMma::FrgTypeA>::value &&
|
||||
cute::is_base_of<cute::UMMA::DescriptorIterator, typename TiledMma::FrgTypeB>::value,
|
||||
"MMA atom must source both A and B operand from smem_desc for this mainloop.");
|
||||
static_assert(
|
||||
(size(AtomThrShapeMNK{}) == 1 &&
|
||||
(cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD> || cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD_MULTICAST>)) ||
|
||||
(size(AtomThrShapeMNK{}) == 2 &&
|
||||
(cute::is_same_v<GmemTiledCopyA, SM100_TMA_2SM_LOAD> || cute::is_same_v<GmemTiledCopyA, SM100_TMA_2SM_LOAD_MULTICAST>)),
|
||||
"GmemTiledCopy - invalid TMA copy atom specified.");
|
||||
static_assert(
|
||||
(size(AtomThrShapeMNK{}) == 1 &&
|
||||
(cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD> || cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_MULTICAST>)) ||
|
||||
(size(AtomThrShapeMNK{}) == 2 &&
|
||||
(cute::is_same_v<GmemTiledCopyB, SM100_TMA_2SM_LOAD> || cute::is_same_v<GmemTiledCopyB, SM100_TMA_2SM_LOAD_MULTICAST>)),
|
||||
"GmemTiledCopy - invalid TMA copy atom specified.");
|
||||
|
||||
using TmaInternalElementA = cute::conditional_t<cute::is_same_v<ElementA, float>, cutlass::tfloat32_t, ElementAMma>;
|
||||
using TmaInternalElementB = cute::conditional_t<cute::is_same_v<ElementB, float>, cutlass::tfloat32_t, ElementBMma>;
|
||||
|
||||
using SmemAllocTypeA = cute::conditional_t<cute::sizeof_bits_v<ElementAMma> < 8, uint8_t, ElementAMma>;
|
||||
using SmemAllocTypeB = cute::conditional_t<cute::sizeof_bits_v<ElementBMma> < 8, uint8_t, ElementBMma>;
|
||||
|
||||
using BitTypeElementA = uint_bit_t<cute::sizeof_bits_v<ElementA>>;
|
||||
using BitTypeElementB = uint_bit_t<cute::sizeof_bits_v<ElementB>>;
|
||||
|
||||
using ArrayElementA = cute::conditional_t<IsRuntimeDataTypeA, BitTypeElementA, ElementA>;
|
||||
using ArrayElementB = cute::conditional_t<IsRuntimeDataTypeB, BitTypeElementB, ElementB>;
|
||||
|
||||
using RuntimeDataTypeA = cute::conditional_t<IsRuntimeDataTypeA, cute::UMMA::MXF8F6F4Format, void*>;
|
||||
using RuntimeDataTypeB = cute::conditional_t<IsRuntimeDataTypeB, cute::UMMA::MXF8F6F4Format, void*>;
|
||||
|
||||
struct SharedStorage {
|
||||
struct TensorStorage : cute::aligned_struct<128, _0> {
|
||||
cute::ArrayEngine<SmemAllocTypeA, cute::cosize_v<SmemLayoutA>> smem_A;
|
||||
cute::ArrayEngine<SmemAllocTypeB, cute::cosize_v<SmemLayoutB>> smem_B;
|
||||
} tensors;
|
||||
|
||||
struct TensorMapStorage : cute::aligned_struct<128, _0> {
|
||||
cute::TmaDescriptor smem_tensormap_A;
|
||||
cute::TmaDescriptor smem_tensormap_B;
|
||||
} tensormaps;
|
||||
|
||||
using PipelineStorage = typename MainloopPipeline::SharedStorage;
|
||||
PipelineStorage pipeline;
|
||||
};
|
||||
|
||||
// Expose shared storage for tensors/pipelines separately to allow kernel layer to reorder them.
|
||||
using TensorStorage = typename SharedStorage::TensorStorage;
|
||||
using TensorMapStorage = typename SharedStorage::TensorMapStorage;
|
||||
using PipelineStorage = typename SharedStorage::PipelineStorage;
|
||||
|
||||
// Only one thread issues the TMA and updates the barriers in a 2SM MMA, adjust bytes accordingly
|
||||
static constexpr uint32_t TmaTransactionBytes =
|
||||
cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutA{})) * cute::sizeof_bits_v<ElementA>) +
|
||||
cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutB{})) * cute::sizeof_bits_v<ElementB>);
|
||||
|
||||
static constexpr bool IsGroupedGemmKernel = !cute::is_same_v<InternalStrideA, StrideA>;
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
ArrayElementA const** ptr_A{nullptr};
|
||||
StrideA dA{};
|
||||
ArrayElementB const** ptr_B{nullptr};
|
||||
StrideB dB{};
|
||||
RuntimeDataTypeA runtime_data_type_a{};
|
||||
RuntimeDataTypeB runtime_data_type_b{};
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
using ClusterLayout_VMNK = decltype(tiled_divide(make_layout(conditional_return<IsDynamicCluster>(make_shape(uint32_t(0), uint32_t(0), Int<1>{}), ClusterShape{})),
|
||||
make_tile(typename TiledMma::AtomThrID{})));
|
||||
|
||||
using TMA_A = decltype(make_tma_atom_A_sm100<TmaInternalElementA>(
|
||||
GmemTiledCopyA{},
|
||||
make_tensor(recast_ptr<TmaInternalElementA>(nullptr), repeat_like(InternalStrideA{}, int32_t(0)), InternalStrideA{}),
|
||||
SmemLayoutA{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
ClusterLayout_VMNK{})
|
||||
);
|
||||
|
||||
using TMA_B = decltype(make_tma_atom_B_sm100<TmaInternalElementB>(
|
||||
GmemTiledCopyB{},
|
||||
make_tensor(recast_ptr<TmaInternalElementB>(nullptr), repeat_like(InternalStrideB{}, int32_t(0)), InternalStrideB{}),
|
||||
SmemLayoutB{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
ClusterLayout_VMNK{})
|
||||
);
|
||||
|
||||
TMA_A tma_load_a;
|
||||
TMA_B tma_load_b;
|
||||
TMA_A tma_load_a_fallback;
|
||||
TMA_B tma_load_b_fallback;
|
||||
dim3 cluster_shape_fallback;
|
||||
RuntimeDataTypeA runtime_data_type_a;
|
||||
RuntimeDataTypeB runtime_data_type_b;
|
||||
cute::TmaDescriptor* tensormaps;
|
||||
ArrayElementA const** ptr_A;
|
||||
StrideA dA;
|
||||
ArrayElementB const** ptr_B;
|
||||
StrideB dB;
|
||||
};
|
||||
|
||||
CUTLASS_DEVICE
|
||||
CollectiveMma(Params const& params, ClusterShape cluster_shape, uint32_t block_rank_in_cluster)
|
||||
: cluster_shape_(cluster_shape)
|
||||
, block_rank_in_cluster_(block_rank_in_cluster) {
|
||||
if constexpr (IsDynamicCluster) {
|
||||
const bool is_fallback_cluster = (cute::size<0>(cluster_shape_) == params.cluster_shape_fallback.x &&
|
||||
cute::size<1>(cluster_shape_) == params.cluster_shape_fallback.y);
|
||||
observed_tma_load_a_ = is_fallback_cluster ? ¶ms.tma_load_a_fallback : ¶ms.tma_load_a;
|
||||
observed_tma_load_b_ = is_fallback_cluster ? ¶ms.tma_load_b_fallback : ¶ms.tma_load_b;
|
||||
}
|
||||
else {
|
||||
observed_tma_load_a_ = ¶ms.tma_load_a;
|
||||
observed_tma_load_b_ = ¶ms.tma_load_b;
|
||||
}
|
||||
}
|
||||
|
||||
template <class ProblemShape>
|
||||
static constexpr Params
|
||||
to_underlying_arguments(
|
||||
ProblemShape problem_shapes,
|
||||
Arguments const& args,
|
||||
void* workspace,
|
||||
cutlass::KernelHardwareInfo const& hw_info = cutlass::KernelHardwareInfo{}) {
|
||||
// These tensor shapes (only applicable for grouped gemm) and pointers are only used to create tensormap/tma desc.
|
||||
// These will be replaced with correct values before the initial tma load.
|
||||
auto init_shape = repeat_like(append<4>(typename ProblemShape::UnderlyingProblemShape{}, 1), int32_t(1));
|
||||
auto init_M = get<0>(init_shape);
|
||||
auto init_N = get<1>(init_shape);
|
||||
auto init_K = get<2>(init_shape);
|
||||
auto init_L = get<3>(init_shape);
|
||||
|
||||
// Tensor pointers will be fixed before the first access
|
||||
TmaInternalElementA const* ptr_A_first_batch = nullptr;
|
||||
TmaInternalElementB const* ptr_B_first_batch = nullptr;
|
||||
|
||||
InternalStrideA stride_a;
|
||||
InternalStrideB stride_b;
|
||||
if constexpr (IsGroupedGemmKernel) {
|
||||
// Strides for Grouped Gemm will be replaced prior to the first access regardless.
|
||||
stride_a = InternalStrideA{};
|
||||
stride_b = InternalStrideB{};
|
||||
}
|
||||
else {
|
||||
// Tensor shapes for Ptr-Array are initialized correctly only here.
|
||||
auto problem_shape_MNK = problem_shapes.get_host_problem_shape(0);
|
||||
init_M = get<0>(problem_shape_MNK);
|
||||
init_N = get<1>(problem_shape_MNK);
|
||||
init_K = get<2>(problem_shape_MNK);
|
||||
|
||||
stride_a = args.dA;
|
||||
stride_b = args.dB;
|
||||
}
|
||||
|
||||
// Batches/Groups are managed by using appropriate pointers to input matrices.
|
||||
Tensor tensor_a = make_tensor(ptr_A_first_batch, make_layout(make_shape(init_M,init_K,init_L), stride_a));
|
||||
Tensor tensor_b = make_tensor(ptr_B_first_batch, make_layout(make_shape(init_N,init_K,init_L), stride_b));
|
||||
|
||||
auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape);
|
||||
// Cluster layout for TMA construction
|
||||
auto cluster_layout_vmnk = tiled_divide(make_layout(cluster_shape), make_tile(typename TiledMma::AtomThrID{}));
|
||||
auto cluster_shape_fallback = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape_fallback);
|
||||
auto cluster_layout_vmnk_fallback = tiled_divide(make_layout(cluster_shape_fallback), make_tile(typename TiledMma::AtomThrID{}));
|
||||
|
||||
typename Params::TMA_A tma_load_a = make_tma_atom_A_sm100<TmaInternalElementA>(
|
||||
GmemTiledCopyA{},
|
||||
tensor_a,
|
||||
SmemLayoutA{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
cluster_layout_vmnk);
|
||||
|
||||
typename Params::TMA_B tma_load_b = make_tma_atom_B_sm100<TmaInternalElementB>(
|
||||
GmemTiledCopyB{},
|
||||
tensor_b,
|
||||
SmemLayoutB{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
cluster_layout_vmnk);
|
||||
|
||||
typename Params::TMA_A tma_load_a_fallback = make_tma_atom_A_sm100<TmaInternalElementA>(
|
||||
GmemTiledCopyA{},
|
||||
tensor_a,
|
||||
SmemLayoutA{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
cluster_layout_vmnk_fallback);
|
||||
|
||||
typename Params::TMA_B tma_load_b_fallback = make_tma_atom_B_sm100<TmaInternalElementB>(
|
||||
GmemTiledCopyB{},
|
||||
tensor_b,
|
||||
SmemLayoutB{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
cluster_layout_vmnk_fallback);
|
||||
|
||||
return {
|
||||
tma_load_a,
|
||||
tma_load_b,
|
||||
tma_load_a_fallback,
|
||||
tma_load_b_fallback,
|
||||
hw_info.cluster_shape_fallback,
|
||||
args.runtime_data_type_a,
|
||||
args.runtime_data_type_b,
|
||||
reinterpret_cast<cute::TmaDescriptor*>(workspace),
|
||||
reinterpret_cast<ArrayElementA const**>(args.ptr_A),
|
||||
args.dA,
|
||||
reinterpret_cast<ArrayElementB const**>(args.ptr_B),
|
||||
args.dB
|
||||
};
|
||||
}
|
||||
|
||||
template <class ProblemShape>
|
||||
static size_t
|
||||
get_workspace_size(ProblemShape const& problem_shape, Arguments const& args, int sm_count) {
|
||||
constexpr uint32_t NumInputTensors = 2;
|
||||
constexpr size_t SizeOfCuTensorMap = sizeof(cute::TmaDescriptor);
|
||||
// Allocate gmem space for input tensormaps per each SM, A tensormap copies followed by B tensormap copies
|
||||
return (NumInputTensors * SizeOfCuTensorMap * sm_count);
|
||||
}
|
||||
|
||||
template <class ProblemShape>
|
||||
static cutlass::Status
|
||||
initialize_workspace(ProblemShape const& problem_shape, Arguments const& args, void* workspace, cudaStream_t stream, CudaHostAdapter* cuda_adapter = nullptr) {
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
|
||||
template<class ProblemShape>
|
||||
static bool
|
||||
can_implement(
|
||||
ProblemShape problem_shapes,
|
||||
[[maybe_unused]] Arguments const& args) {
|
||||
static constexpr bool IsF8F6F4 = detail::is_sm100_mma_f8f6f4<TiledMma, ElementA, ElementB>();
|
||||
constexpr int tma_alignment_bits_A = cutlass::detail::get_input_alignment_bits<ElementA, IsF8F6F4>();
|
||||
constexpr int tma_alignment_bits_B = cutlass::detail::get_input_alignment_bits<ElementB, IsF8F6F4>();
|
||||
constexpr int min_tma_aligned_elements_A = tma_alignment_bits_A / cute::sizeof_bits<ElementA>::value;
|
||||
constexpr int min_tma_aligned_elements_B = tma_alignment_bits_B / cute::sizeof_bits<ElementB>::value;
|
||||
|
||||
bool implementable = true;
|
||||
if (problem_shapes.is_host_problem_shape_available()) {
|
||||
// Check alignment for all problem sizes
|
||||
for (int i = 0; i < problem_shapes.groups(); i++) {
|
||||
auto problem_shape_MNKL = append<4>(problem_shapes.get_host_problem_shape(i), 1);
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
implementable = implementable && cutlass::detail::check_alignment<min_tma_aligned_elements_A>(cute::make_shape(M,K,L), InternalStrideA{});
|
||||
implementable = implementable && cutlass::detail::check_alignment<min_tma_aligned_elements_B>(cute::make_shape(N,K,L), InternalStrideB{});
|
||||
}
|
||||
}
|
||||
|
||||
if (!implementable) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Problem Size doesn't meet the minimum alignment requirements for TMA.\n");
|
||||
}
|
||||
return implementable;
|
||||
}
|
||||
|
||||
/// Construct A Single Stage's Accumulator Shape
|
||||
CUTLASS_DEVICE auto
|
||||
partition_accumulator_shape() {
|
||||
auto acc_shape = partition_shape_C(TiledMma{}, take<0,2>(TileShape{})); // ((MMA_TILE_M,MMA_TILE_N),MMA_M,MMA_N)
|
||||
|
||||
return acc_shape;
|
||||
}
|
||||
|
||||
template <class FrgEngine, class FrgLayout>
|
||||
CUTLASS_DEVICE auto
|
||||
slice_accumulator(cute::Tensor<FrgEngine, FrgLayout> const& accumulators, int stage) {
|
||||
return accumulators(_,_,_,stage);
|
||||
}
|
||||
|
||||
/// Set up the data needed by this collective for load.
|
||||
/// Return tuple element contain
|
||||
/// gA_mkl - The tiled tma tensor for input A
|
||||
/// gB_nkl - The tiled tma tensor for input B
|
||||
/// tAsA - partitioned smem tensor for A
|
||||
/// tBsB - partitioned smem tensor for B
|
||||
/// mcast_mask_a - tma multicast mask for A
|
||||
/// mcast_mask_b - tma multicast mask for B
|
||||
template <class ProblemShape_MNKL>
|
||||
CUTLASS_DEVICE auto
|
||||
load_init(
|
||||
ProblemShape_MNKL const& problem_shape_MNKL,
|
||||
Params const& params,
|
||||
TensorStorage& shared_tensors,
|
||||
TensorMapStorage& shared_tensormaps,
|
||||
int32_t const sm_count, int32_t const sm_idx,
|
||||
[[maybe_unused]] int32_t init_group) const {
|
||||
using X = Underscore;
|
||||
|
||||
// Separate out problem shape for convenience
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
// Problem Shape and therefore strides that we construct are [M,N,K,L], but since here for the TMA loads
|
||||
// we are managing TMA descriptors to change batches, we need to neglect the L mode
|
||||
const int32_t mock_L = 1;
|
||||
|
||||
// Represent the full tensors -- get these from TMA
|
||||
Tensor mA_mkl = observed_tma_load_a_->get_tma_tensor(make_shape(M,K,mock_L));
|
||||
Tensor mB_nkl = observed_tma_load_b_->get_tma_tensor(make_shape(N,K,mock_L));
|
||||
|
||||
// Tile the tensors and defer the slice
|
||||
Tensor gA_mkl = local_tile(mA_mkl, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M, BLK_K, m, k, l)
|
||||
Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N, BLK_K, n, k, l)
|
||||
|
||||
// Partition for this CTA
|
||||
ThrMMA cta_mma = TiledMma{}.get_slice(blockIdx.x % size(typename TiledMma::AtomThrID{}));
|
||||
|
||||
Tensor tCgA_mkl = cta_mma.partition_A(gA_mkl); // (MMA, MMA_M, MMA_K, m, k, l)
|
||||
Tensor tCgB_nkl = cta_mma.partition_B(gB_nkl); // (MMA, MMA_N, MMA_K, n, k, l)
|
||||
|
||||
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (MMA,MMA_M,MMA_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (MMA,MMA_N,MMA_K,PIPE)
|
||||
|
||||
// Define the CTA-in-Cluster Layout and Coord
|
||||
Layout cta_layout_mnk = make_layout(cluster_shape_);
|
||||
Layout cta_layout_vmnk = tiled_divide(cta_layout_mnk, make_tile(typename TiledMma::AtomThrID{}));
|
||||
auto cta_coord_vmnk = cta_layout_vmnk.get_flat_coord(block_rank_in_cluster_);
|
||||
|
||||
// Project the cta_layout for tma_a along the n-modes
|
||||
auto [tAgA_mkl, tAsA] = tma_partition(*observed_tma_load_a_,
|
||||
get<2>(cta_coord_vmnk), make_layout(size<2>(cta_layout_vmnk)),
|
||||
group_modes<0,3>(sA), group_modes<0,3>(tCgA_mkl));
|
||||
|
||||
// Project the cta_layout for tma_b along the m-modes
|
||||
auto [tBgB_nkl, tBsB] = tma_partition(*observed_tma_load_b_,
|
||||
get<1>(cta_coord_vmnk), make_layout(size<1>(cta_layout_vmnk)),
|
||||
group_modes<0,3>(sB), group_modes<0,3>(tCgB_nkl));
|
||||
|
||||
// TMA Multicast Masks
|
||||
uint16_t mcast_mask_a = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk);
|
||||
uint16_t mcast_mask_b = create_tma_multicast_mask<1>(cta_layout_vmnk, cta_coord_vmnk);
|
||||
|
||||
// Fetch a copy of tensormaps for the CTA from Params
|
||||
auto input_tensormaps = tensormaps_init(params, shared_tensormaps, sm_count, sm_idx);
|
||||
|
||||
return cute::make_tuple(
|
||||
gA_mkl, gB_nkl, // for scheduler
|
||||
tAgA_mkl, tBgB_nkl, tAsA, tBsB, // for input tensor values
|
||||
mcast_mask_a, mcast_mask_b, // multicast masks
|
||||
input_tensormaps); // for tma descriptor modification (per-CTA tensormap copy)
|
||||
}
|
||||
|
||||
/// Set up the data needed by this collective for mma compute.
|
||||
template <class FrgEngine, class FrgLayout>
|
||||
CUTLASS_DEVICE auto
|
||||
mma_init(
|
||||
Params const& params,
|
||||
[[maybe_unused]] cute::Tensor<FrgEngine, FrgLayout> const& accumulators,
|
||||
TensorStorage& shared_tensors,
|
||||
[[maybe_unused]] uint32_t const tmem_nonaccum_offset) const {
|
||||
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
|
||||
|
||||
// Allocate "fragments/descriptors" for A and B matrices
|
||||
Tensor tCrA = TiledMma::make_fragment_A(sA); // (MMA,MMA_M,MMA_K,PIPE)
|
||||
Tensor tCrB = TiledMma::make_fragment_B(sB); // (MMA,MMA_N,MMA_K,PIPE)
|
||||
|
||||
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sA)); // PIPE
|
||||
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sB));
|
||||
|
||||
TiledMma tiled_mma;
|
||||
|
||||
if constexpr (IsRuntimeDataType) {
|
||||
// Update instruction descriptor according to runtime argument.
|
||||
// Applying bitmask (0b111) to help compiler deduce that the conversion and assignment are safe.
|
||||
tiled_mma.idesc_.a_format_ = uint8_t(params.runtime_data_type_a) & 0b111;
|
||||
tiled_mma.idesc_.b_format_ = uint8_t(params.runtime_data_type_b) & 0b111;
|
||||
}
|
||||
|
||||
return cute::make_tuple(tiled_mma, tCrA, tCrB);
|
||||
}
|
||||
|
||||
/// Perform a collective-scoped matrix multiply-accumulate
|
||||
/// Producer Perspective
|
||||
template <
|
||||
class GTensorA, class GTensorB,
|
||||
class GTensorPartitionedA, class GTensorPartitionedB,
|
||||
class STensorA, class STensorB,
|
||||
class TensorMapA, class TensorMapB,
|
||||
class TileCoordMNKL,
|
||||
class KTileIterator
|
||||
>
|
||||
CUTLASS_DEVICE auto
|
||||
load(
|
||||
Params const& params,
|
||||
MainloopPipeline mainloop_pipeline,
|
||||
MainloopPipelineState mainloop_pipe_producer_state,
|
||||
cute::tuple<GTensorA, GTensorB,
|
||||
GTensorPartitionedA, GTensorPartitionedB,
|
||||
STensorA, STensorB,
|
||||
uint16_t, uint16_t,
|
||||
cute::tuple<TensorMapA, TensorMapB>> const& load_inputs,
|
||||
TileCoordMNKL const& cta_coord_mnkl,
|
||||
KTileIterator k_tile_iter, int k_tile_count,
|
||||
bool did_batch_change) {
|
||||
|
||||
auto [unused_gA, unused_gB,
|
||||
tAgA_mkl, tBgB_nkl, tAsA, tBsB,
|
||||
mcast_mask_a, mcast_mask_b,
|
||||
input_tensormaps] = load_inputs;
|
||||
|
||||
// Check to see if tensormaps have been replaced in gmem
|
||||
if (did_batch_change) {
|
||||
tensormaps_fence_acquire(input_tensormaps);
|
||||
}
|
||||
|
||||
// slice out the work coord from partitioned tensors
|
||||
Tensor tAgA = tAgA_mkl(_, get<0>(cta_coord_mnkl) / size(typename TiledMma::AtomThrID{}), _, get<3>(cta_coord_mnkl));
|
||||
Tensor tBgB = tBgB_nkl(_, get<1>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl));
|
||||
|
||||
auto barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
|
||||
|
||||
// Issue the Mainloop loads
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
while (k_tile_count > 0) {
|
||||
// LOCK mainloop_pipe_producer_state for _writing_
|
||||
mainloop_pipeline.producer_acquire(mainloop_pipe_producer_state, barrier_token);
|
||||
|
||||
using BarrierType = typename MainloopPipeline::ProducerBarrierType;
|
||||
BarrierType* tma_barrier = mainloop_pipeline.producer_get_barrier(mainloop_pipe_producer_state);
|
||||
|
||||
int write_stage = mainloop_pipe_producer_state.index();
|
||||
++mainloop_pipe_producer_state;
|
||||
barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
|
||||
|
||||
if (cute::elect_one_sync()) {
|
||||
copy(observed_tma_load_a_->with(get<0>(input_tensormaps), *tma_barrier, mcast_mask_a), tAgA(_,*k_tile_iter), tAsA(_,write_stage));
|
||||
copy(observed_tma_load_b_->with(get<1>(input_tensormaps), *tma_barrier, mcast_mask_b), tBgB(_,*k_tile_iter), tBsB(_,write_stage));
|
||||
}
|
||||
--k_tile_count;
|
||||
++k_tile_iter;
|
||||
}
|
||||
|
||||
return cute::make_tuple(mainloop_pipe_producer_state, k_tile_iter);
|
||||
}
|
||||
|
||||
/// Perform a Producer Epilogue to prevent early exit of ctas in a Cluster
|
||||
CUTLASS_DEVICE void
|
||||
load_tail(MainloopPipeline mainloop_pipeline, MainloopPipelineState mainloop_pipe_producer_state) {
|
||||
// Issue the epilogue waits
|
||||
// This helps avoid early exit of ctas in Cluster
|
||||
// Waits for all stages to either be released (all
|
||||
// Consumer UNLOCKs), or if the stage was never used
|
||||
// then would just be acquired since the phase was
|
||||
// still inverted from make_producer_start_state
|
||||
mainloop_pipeline.producer_tail(mainloop_pipe_producer_state);
|
||||
}
|
||||
|
||||
/// Perform a collective-scoped matrix multiply-accumulate
|
||||
/// Consumer Perspective
|
||||
template <
|
||||
class AccumulatorPipeline,
|
||||
class FrgEngine, class FrgLayout,
|
||||
class FragmentA, class FragmentB,
|
||||
class CtaTileCoord
|
||||
>
|
||||
CUTLASS_DEVICE auto
|
||||
mma(cute::tuple<MainloopPipeline,
|
||||
AccumulatorPipeline> pipelines,
|
||||
cute::tuple<MainloopPipelineState,
|
||||
typename AccumulatorPipeline::PipelineState> pipeline_states,
|
||||
cute::Tensor<FrgEngine, FrgLayout>& accumulators,
|
||||
cute::tuple<TiledMma, FragmentA, FragmentB> const& mma_inputs,
|
||||
CtaTileCoord cta_tile_coord,
|
||||
int k_tile_count
|
||||
) {
|
||||
static_assert(is_tmem<FrgEngine>::value, "Accumulator must be tmem resident.");
|
||||
static_assert(rank(FrgLayout{}) == 3, "Accumulator must be MMA-partitioned: (MMA, MMA_M, MMA_N)");
|
||||
auto [tiled_mma, tCrA, tCrB] = mma_inputs;
|
||||
|
||||
auto [mainloop_pipeline, accumulator_pipeline] = pipelines;
|
||||
auto [mainloop_pipe_consumer_state, accumulator_pipe_producer_state] = pipeline_states;
|
||||
|
||||
uint32_t skip_wait = k_tile_count <= 0;
|
||||
auto barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
|
||||
|
||||
//
|
||||
// PIPELINED MAIN LOOP
|
||||
//
|
||||
tiled_mma.accumulate_ = UMMA::ScaleOut::Zero;
|
||||
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
while (k_tile_count > 0) {
|
||||
// WAIT on mainloop_pipe_consumer_state until its data are available
|
||||
// (phase bit flips from mainloop_pipe_consumer_state.phase() value)
|
||||
mainloop_pipeline.consumer_wait(mainloop_pipe_consumer_state, barrier_token);
|
||||
|
||||
// Compute on k_tile
|
||||
int read_stage = mainloop_pipe_consumer_state.index();
|
||||
// Save current mainlop pipeline read state
|
||||
auto curr_mainloop_pipe_consumer_state = mainloop_pipe_consumer_state;
|
||||
|
||||
// Advance mainloop_pipe
|
||||
++mainloop_pipe_consumer_state;
|
||||
--k_tile_count;
|
||||
skip_wait = k_tile_count <= 0;
|
||||
// Peek at next iteration
|
||||
barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
|
||||
|
||||
// Unroll the K mode manually so we can set scale C to 1
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
|
||||
// (V,M) x (V,N) => (V,M,N)
|
||||
cute::gemm(tiled_mma,
|
||||
tCrA(_,_,k_block,read_stage),
|
||||
tCrB(_,_,k_block,read_stage),
|
||||
accumulators);
|
||||
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
|
||||
}
|
||||
mainloop_pipeline.consumer_release(curr_mainloop_pipe_consumer_state);
|
||||
}
|
||||
|
||||
return mainloop_pipe_consumer_state;
|
||||
}
|
||||
|
||||
//
|
||||
// Methods to perform different parts of TMA/Tensormap modifications
|
||||
//
|
||||
|
||||
CUTLASS_DEVICE auto
|
||||
tensormaps_init(
|
||||
Params const& mainloop_params,
|
||||
TensorMapStorage& shared_tensormaps,
|
||||
int32_t const sm_count,
|
||||
int32_t const sm_idx) const {
|
||||
cute::TmaDescriptor* gmem_tensormap = mainloop_params.tensormaps;
|
||||
|
||||
cute::TmaDescriptor* tma_desc_a = &gmem_tensormap[sm_idx];
|
||||
cute::TmaDescriptor* tma_desc_b = &gmem_tensormap[sm_idx + sm_count];
|
||||
|
||||
if (cute::elect_one_sync()) {
|
||||
// Bringing tensormaps from params to smem for modification later
|
||||
Tensor pA_tensormap = make_tensor(observed_tma_load_a_->get_tma_descriptor(), Int<1>{}, Int<1>{});
|
||||
Tensor sA_tensormap = make_tensor(make_smem_ptr(&shared_tensormaps.smem_tensormap_A), Int<1>{}, Int<1>{});
|
||||
Tensor pB_tensormap = make_tensor(observed_tma_load_b_->get_tma_descriptor(), Int<1>{}, Int<1>{});
|
||||
Tensor sB_tensormap = make_tensor(make_smem_ptr(&shared_tensormaps.smem_tensormap_B), Int<1>{}, Int<1>{});
|
||||
|
||||
copy(recast<uint128_t>(pA_tensormap), recast<uint128_t>(sA_tensormap));
|
||||
copy(recast<uint128_t>(pB_tensormap), recast<uint128_t>(sB_tensormap));
|
||||
}
|
||||
__syncwarp();
|
||||
|
||||
return cute::make_tuple(tma_desc_a, tma_desc_b);
|
||||
}
|
||||
|
||||
// Replace address for the global tensor (to be done by single thread)
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_replace_global_address(
|
||||
TensorMapStorage& shared_tensormaps,
|
||||
Params const& mainloop_params,
|
||||
int32_t next_batch) {
|
||||
// Replacing global_address for the next batch
|
||||
cute::tma_descriptor_replace_addr_in_shared_mem(shared_tensormaps.smem_tensormap_A,
|
||||
mainloop_params.ptr_A[next_batch]);
|
||||
cute::tma_descriptor_replace_addr_in_shared_mem(shared_tensormaps.smem_tensormap_B,
|
||||
mainloop_params.ptr_B[next_batch]);
|
||||
}
|
||||
|
||||
// Replace dim and strides for the global tensor - used only for Grouped GEMM (to be done by single thread)
|
||||
template <class ProblemShape_MNKL>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_replace_global_tensor_properties(
|
||||
TensorMapStorage& shared_tensormaps,
|
||||
Params const& mainloop_params,
|
||||
int32_t next_group,
|
||||
ProblemShape_MNKL problem_shape_mnkl) {
|
||||
const uint32_t M = get<0>(problem_shape_mnkl);
|
||||
const uint32_t N = get<1>(problem_shape_mnkl);
|
||||
const uint32_t K = get<2>(problem_shape_mnkl);
|
||||
// Replace all dims for consistency
|
||||
constexpr int MaxTensorRank = 5;
|
||||
cute::array<uint32_t, MaxTensorRank> prob_shape_A = {1,1,1,1,1};
|
||||
cute::array<uint64_t, MaxTensorRank> prob_stride_A = {0,0,0,0,0};
|
||||
cute::array<uint32_t, MaxTensorRank> prob_shape_B = {1,1,1,1,1};
|
||||
cute::array<uint64_t, MaxTensorRank> prob_stride_B = {0,0,0,0,0};
|
||||
|
||||
TmaInternalElementA const* ptr_A = nullptr;
|
||||
Tensor tensor_a = make_tensor(ptr_A, make_shape(M,K,Int<1>{}), mainloop_params.dA[next_group]);
|
||||
|
||||
TmaInternalElementB const* ptr_B = nullptr;
|
||||
Tensor tensor_b = make_tensor(ptr_B, make_shape(N,K,Int<1>{}), mainloop_params.dB[next_group]);
|
||||
|
||||
cute::detail::fill_tma_gmem_shape_stride(*observed_tma_load_a_, tensor_a,
|
||||
prob_shape_A, prob_stride_A);
|
||||
cute::detail::fill_tma_gmem_shape_stride(*observed_tma_load_b_, tensor_b,
|
||||
prob_shape_B, prob_stride_B);
|
||||
|
||||
// Convert strides to byte strides
|
||||
for (uint64_t& stride : prob_stride_A) {
|
||||
stride = (stride * sizeof_bits_v<TmaInternalElementA>) / 8;
|
||||
}
|
||||
for (uint64_t& stride : prob_stride_B) {
|
||||
stride = (stride * sizeof_bits_v<TmaInternalElementB>) / 8;
|
||||
}
|
||||
|
||||
cute::tma_descriptor_replace_dims_strides_in_shared_mem(shared_tensormaps.smem_tensormap_A,
|
||||
prob_shape_A,
|
||||
prob_stride_A);
|
||||
cute::tma_descriptor_replace_dims_strides_in_shared_mem(shared_tensormaps.smem_tensormap_B,
|
||||
prob_shape_B,
|
||||
prob_stride_B);
|
||||
}
|
||||
|
||||
// The entire warp must call this function collectively (that is, the instructions are aligned)
|
||||
template <class TensorMapA, class TensorMapB, class ProblemShape>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_perform_update(
|
||||
TensorMapStorage& shared_tensormaps,
|
||||
Params const& mainloop_params,
|
||||
cute::tuple<TensorMapA, TensorMapB> const& input_tensormaps,
|
||||
ProblemShape problem_shape,
|
||||
int32_t next_batch) {
|
||||
if (cute::elect_one_sync()) {
|
||||
// Replacing global_address for the next batch
|
||||
tensormaps_replace_global_address(shared_tensormaps, mainloop_params, next_batch);
|
||||
|
||||
if constexpr (IsGroupedGemmKernel) {
|
||||
auto problem_shape_MNKL = append<4>(problem_shape.get_problem_shape(next_batch), 1);
|
||||
// Replacing global dims and strides for the next batch
|
||||
tensormaps_replace_global_tensor_properties(shared_tensormaps,
|
||||
mainloop_params, next_batch, problem_shape_MNKL);
|
||||
}
|
||||
}
|
||||
// Ensure warp is converged before issuing tensormap fence release
|
||||
__syncwarp();
|
||||
// Entire warp must do this (ie its aligned)
|
||||
tensormaps_cp_fence_release(shared_tensormaps, input_tensormaps);
|
||||
}
|
||||
|
||||
template <class TensorMapA, class TensorMapB>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_cp_fence_release (
|
||||
TensorMapStorage& shared_tensormaps,
|
||||
cute::tuple<TensorMapA, TensorMapB> const& input_tensormaps) {
|
||||
if (cute::elect_one_sync()) {
|
||||
cute::tma_desc_commit_group();
|
||||
cute::tma_desc_wait_group();
|
||||
}
|
||||
// Entire warp must do this (i.e. it's aligned)
|
||||
tma_descriptor_cp_fence_release(get<0>(input_tensormaps), shared_tensormaps.smem_tensormap_A);
|
||||
tma_descriptor_cp_fence_release(get<1>(input_tensormaps), shared_tensormaps.smem_tensormap_B);
|
||||
}
|
||||
|
||||
// The entire warp must call this function collectively (that is, the instructions are aligned)
|
||||
template <class TensorMapA, class TensorMapB>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
tensormaps_fence_acquire(cute::tuple<TensorMapA, TensorMapB> const& input_tensormaps) {
|
||||
cute::tma_descriptor_fence_acquire(get<0>(input_tensormaps));
|
||||
cute::tma_descriptor_fence_acquire(get<1>(input_tensormaps));
|
||||
}
|
||||
|
||||
private:
|
||||
|
||||
typename Params::TMA_A const* observed_tma_load_a_{nullptr};
|
||||
typename Params::TMA_B const* observed_tma_load_b_{nullptr};
|
||||
|
||||
ClusterShape cluster_shape_;
|
||||
uint32_t block_rank_in_cluster_;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -0,0 +1,723 @@
|
||||
/***************************************************************************************************
|
||||
* 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/detail/collective.hpp"
|
||||
#include "cutlass/detail/cluster.hpp"
|
||||
#include "cutlass/gemm/dispatch_policy.hpp"
|
||||
#include "cutlass/numeric_types.h"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/gemm/gemm.h"
|
||||
#include "cutlass/trace.h"
|
||||
#include "cutlass/kernel_hardware_info.hpp"
|
||||
#include "cutlass/detail/sm100_tmem_helper.hpp"
|
||||
|
||||
#include "cute/algorithm/functional.hpp"
|
||||
#include "cute/arch/cluster_sm90.hpp"
|
||||
#include "cute/atom/mma_atom.hpp"
|
||||
#include "cute/algorithm/gemm.hpp"
|
||||
#include "cute/tensor_predicate.hpp"
|
||||
#include "cute/numeric/arithmetic_tuple.hpp"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::collective {
|
||||
using namespace cute;
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// WarpSpecialized Mainloop
|
||||
// Both DMA Load and MMA methods of this class must be run by a single thread that's picked by elect_one
|
||||
template <
|
||||
int Stages,
|
||||
int SchedulerPipelineStageCount,
|
||||
int AccumulatorPipelineStageCount,
|
||||
class ClusterShape, // Static cluster shape or dynamic (int, int, _1)
|
||||
class TileShape_, // (MmaAtomShapeM, MmaAtomShapeN, TileK)
|
||||
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<
|
||||
MainloopSm100TmaUmmaWarpSpecialized<
|
||||
Stages,
|
||||
SchedulerPipelineStageCount,
|
||||
AccumulatorPipelineStageCount,
|
||||
ClusterShape>,
|
||||
TileShape_,
|
||||
ElementA_,
|
||||
StrideA_,
|
||||
ElementB_,
|
||||
StrideB_,
|
||||
TiledMma_,
|
||||
GmemTiledCopyA_,
|
||||
SmemLayoutAtomA_,
|
||||
SmemCopyAtomA_,
|
||||
TransformA_,
|
||||
GmemTiledCopyB_,
|
||||
SmemLayoutAtomB_,
|
||||
SmemCopyAtomB_,
|
||||
TransformB_>
|
||||
{
|
||||
//
|
||||
// Type Aliases
|
||||
//
|
||||
using TiledMma = TiledMma_;
|
||||
using AtomThrShapeMNK = Shape<decltype(shape<0>(typename TiledMma::ThrLayoutVMNK{})), _1, _1>;
|
||||
|
||||
using DispatchPolicy = MainloopSm100TmaUmmaWarpSpecialized<
|
||||
Stages,
|
||||
SchedulerPipelineStageCount,
|
||||
AccumulatorPipelineStageCount,
|
||||
ClusterShape>;
|
||||
using TileShape = TileShape_;
|
||||
|
||||
static constexpr bool IsDynamicCluster = not cute::is_static_v<ClusterShape>;
|
||||
|
||||
CUTE_STATIC_ASSERT_V(evenly_divides(TileShape{}, tile_shape(TiledMma{})),
|
||||
"Static cluster shape used: TileShape should be evenly divided by TiledMma");
|
||||
|
||||
using CtaShape_MNK = decltype(shape_div(TileShape{}, AtomThrShapeMNK{}));
|
||||
|
||||
// Define A and B block shapes for reduced size TMA_LOADs
|
||||
using MmaShapeA_MK = decltype(partition_shape_A(TiledMma{}, make_shape(size<0>(TileShape{}), size<2>(TileShape{}))));
|
||||
using MmaShapeB_NK = decltype(partition_shape_B(TiledMma{}, make_shape(size<1>(TileShape{}), size<2>(TileShape{}))));
|
||||
|
||||
using ElementA = ElementA_;
|
||||
using ElementAMma = typename TiledMma::ValTypeA;
|
||||
using StrideA = StrideA_;
|
||||
using ElementB = ElementB_;
|
||||
using ElementBMma = typename TiledMma::ValTypeB;
|
||||
using StrideB = StrideB_;
|
||||
|
||||
static constexpr bool IsRuntimeDataTypeA = cutlass::gemm::collective::detail::is_sm10x_runtime_f8f6f4<ElementA>();
|
||||
|
||||
static constexpr bool IsRuntimeDataTypeB = cutlass::gemm::collective::detail::is_sm10x_runtime_f8f6f4<ElementB>();
|
||||
|
||||
static_assert((IsRuntimeDataTypeA && IsRuntimeDataTypeB) ||
|
||||
(!IsRuntimeDataTypeA && !IsRuntimeDataTypeB),
|
||||
"ElementA and ElementB should be both runtime or both static.");
|
||||
|
||||
static constexpr bool IsRuntimeDataType = IsRuntimeDataTypeA && IsRuntimeDataTypeB;
|
||||
|
||||
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 MainloopPipeline = cutlass::PipelineTmaUmmaAsync<
|
||||
DispatchPolicy::Stages,
|
||||
ClusterShape,
|
||||
AtomThrShapeMNK>;
|
||||
using MainloopPipelineState = typename MainloopPipeline::PipelineState;
|
||||
|
||||
static_assert(rank(SmemLayoutAtomA{}) == 2, "SmemLayoutAtomA must be rank 2 (M,K)");
|
||||
static_assert(((size<0,0>(MmaShapeA_MK{}) * size<1>(MmaShapeA_MK{})) % size<0>(SmemLayoutAtomA{})) == 0,
|
||||
"SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert(((size<0,1>(MmaShapeA_MK{}) * size<2>(MmaShapeA_MK{})) % size<1>(SmemLayoutAtomA{})) == 0,
|
||||
"SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert(cute::is_void_v<SmemCopyAtomA>,
|
||||
"SM100 UMMA cannot have a non-void copy atom for smem sourced instructions.");
|
||||
|
||||
static_assert(rank(SmemLayoutAtomB{}) == 2, "SmemLayoutAtomB must be rank 2 (N,K)");
|
||||
static_assert(((size<0,0>(MmaShapeB_NK{}) * size<1>(MmaShapeB_NK{})) % size<0>(SmemLayoutAtomB{})) == 0,
|
||||
"SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert(((size<0,1>(MmaShapeB_NK{}) * size<2>(MmaShapeB_NK{})) % size<1>(SmemLayoutAtomB{})) == 0,
|
||||
"SmemLayoutAtom must evenly divide tile shape.");
|
||||
static_assert(cute::is_void_v<SmemCopyAtomB>,
|
||||
"SM100 UMMA cannot have a non-void copy atom for smem sourced instructions.");
|
||||
|
||||
// Tile along K mode first before tiling over MN. PIPE mode last as usual.
|
||||
// This maximizes TMA boxes due to better smem-K vectorization, reducing total issued TMAs.
|
||||
// (MMA_TILE_M,MMA_TILE_K),MMA_M,MMA_K,PIPE)
|
||||
using SmemLayoutA = decltype(UMMA::tile_to_mma_shape(
|
||||
SmemLayoutAtomA{},
|
||||
append(MmaShapeA_MK{}, Int<DispatchPolicy::Stages>{}),
|
||||
cute::conditional_t<cutlass::gemm::detail::is_mn_major<StrideA>(), Step<_2,_1,_3>, Step<_1,_2,_3>>{}));
|
||||
// (MMA_TILE_N,MMA_TILE_K),MMA_N,MMA_K,PIPE)
|
||||
using SmemLayoutB = decltype(UMMA::tile_to_mma_shape(
|
||||
SmemLayoutAtomB{},
|
||||
append(MmaShapeB_NK{}, Int<DispatchPolicy::Stages>{}),
|
||||
cute::conditional_t<cutlass::gemm::detail::is_mn_major<StrideB>(), Step<_2,_1,_3>, Step<_1,_2,_3>>{}));
|
||||
|
||||
static_assert(DispatchPolicy::Stages >= 2, "Specialization requires Stages set to value 1 or more.");
|
||||
static_assert(cute::is_base_of<cute::UMMA::DescriptorIterator, typename TiledMma::FrgTypeA>::value &&
|
||||
cute::is_base_of<cute::UMMA::DescriptorIterator, typename TiledMma::FrgTypeB>::value,
|
||||
"MMA atom must source both A and B operand from smem_desc for this mainloop.");
|
||||
static_assert(
|
||||
(size(AtomThrShapeMNK{}) == 1 &&
|
||||
(cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD> || cute::is_same_v<GmemTiledCopyA, SM90_TMA_LOAD_MULTICAST>)) ||
|
||||
(size(AtomThrShapeMNK{}) == 2 &&
|
||||
(cute::is_same_v<GmemTiledCopyA, SM100_TMA_2SM_LOAD> || cute::is_same_v<GmemTiledCopyA, SM100_TMA_2SM_LOAD_MULTICAST>)),
|
||||
"GmemTiledCopy - invalid TMA copy atom specified.");
|
||||
static_assert(
|
||||
(size(AtomThrShapeMNK{}) == 1 &&
|
||||
(cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD> || cute::is_same_v<GmemTiledCopyB, SM90_TMA_LOAD_MULTICAST>)) ||
|
||||
(size(AtomThrShapeMNK{}) == 2 &&
|
||||
(cute::is_same_v<GmemTiledCopyB, SM100_TMA_2SM_LOAD> || cute::is_same_v<GmemTiledCopyB, SM100_TMA_2SM_LOAD_MULTICAST>)),
|
||||
"GmemTiledCopy - invalid TMA copy atom specified.");
|
||||
|
||||
using TmaInternalElementA = cute::conditional_t<cute::is_same_v<ElementA, float>, cutlass::tfloat32_t, ElementAMma>;
|
||||
using TmaInternalElementB = cute::conditional_t<cute::is_same_v<ElementB, float>, cutlass::tfloat32_t, ElementBMma>;
|
||||
|
||||
using SmemAllocTypeA = cute::conditional_t<cute::sizeof_bits_v<ElementAMma> < 8, uint8_t, ElementAMma>;
|
||||
using SmemAllocTypeB = cute::conditional_t<cute::sizeof_bits_v<ElementBMma> < 8, uint8_t, ElementBMma>;
|
||||
|
||||
using BitTypeElementA = cute::uint_bit_t<cute::sizeof_bits_v<ElementA>>;
|
||||
using BitTypeElementB = cute::uint_bit_t<cute::sizeof_bits_v<ElementB>>;
|
||||
|
||||
using ArrayElementA = cute::conditional_t<IsRuntimeDataTypeA, BitTypeElementA, ElementA>;
|
||||
using ArrayElementB = cute::conditional_t<IsRuntimeDataTypeB, BitTypeElementB, ElementB>;
|
||||
|
||||
using RuntimeDataTypeA = cute::conditional_t<IsRuntimeDataTypeA, cute::UMMA::MXF8F6F4Format, void*>;
|
||||
using RuntimeDataTypeB = cute::conditional_t<IsRuntimeDataTypeB, cute::UMMA::MXF8F6F4Format, void*>;
|
||||
|
||||
struct SharedStorage {
|
||||
struct TensorStorage : cute::aligned_struct<128, _0> {
|
||||
cute::ArrayEngine<SmemAllocTypeA, cute::cosize_v<SmemLayoutA>> smem_A;
|
||||
cute::ArrayEngine<SmemAllocTypeB, cute::cosize_v<SmemLayoutB>> smem_B;
|
||||
} tensors;
|
||||
|
||||
using PipelineStorage = typename MainloopPipeline::SharedStorage;
|
||||
PipelineStorage pipeline;
|
||||
};
|
||||
|
||||
// Expose shared storage for tensors/pipelines separately to allow kernel layer to reorder them.
|
||||
using TensorStorage = typename SharedStorage::TensorStorage;
|
||||
using PipelineStorage = typename SharedStorage::PipelineStorage;
|
||||
|
||||
// Only one thread issues the TMA and updates the barriers in a 2SM MMA, adjust bytes accordingly
|
||||
static constexpr uint32_t TmaTransactionBytes =
|
||||
cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutA{})) * cute::sizeof_bits_v<ElementA>) +
|
||||
cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutB{})) * cute::sizeof_bits_v<ElementB>);
|
||||
|
||||
template<class AccTensor>
|
||||
struct TmemStorage {
|
||||
AccTensor accumulators;
|
||||
};
|
||||
|
||||
template<
|
||||
class KTileCount,
|
||||
class GTensorPartitionedA, class GTensorPartitionedB,
|
||||
class STensorA, class STensorB
|
||||
>
|
||||
struct LoadParams {
|
||||
// for scheduler
|
||||
KTileCount k_tiles;
|
||||
// for input tensor values
|
||||
GTensorPartitionedA tAgA_mkl;
|
||||
GTensorPartitionedB tBgB_nkl;
|
||||
STensorA tAsA;
|
||||
STensorB tBsB;
|
||||
// the TMA multicast masks
|
||||
uint16_t mcast_mask_a;
|
||||
uint16_t mcast_mask_b;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
LoadParams (
|
||||
KTileCount k_tiles_,
|
||||
GTensorPartitionedA tAgA_mkl_, GTensorPartitionedB tBgB_nkl_,
|
||||
STensorA tAsA_, STensorB tBsB_,
|
||||
uint16_t mcast_mask_a_, uint16_t mcast_mask_b_)
|
||||
: k_tiles(k_tiles_)
|
||||
, tAgA_mkl(tAgA_mkl_), tBgB_nkl(tBgB_nkl_)
|
||||
, tAsA(tAsA_), tBsB(tBsB_)
|
||||
, mcast_mask_a(mcast_mask_a_), mcast_mask_b(mcast_mask_b_) {}
|
||||
};
|
||||
|
||||
template<class FragmentA, class FragmentB>
|
||||
struct MmaParams {
|
||||
TiledMma tiled_mma;
|
||||
FragmentA tCrA;
|
||||
FragmentB tCrB;
|
||||
|
||||
CUTLASS_DEVICE
|
||||
MmaParams (
|
||||
TiledMma tiled_mma_,
|
||||
FragmentA tCrA_, FragmentB tCrB_)
|
||||
: tiled_mma(tiled_mma_)
|
||||
, tCrA(tCrA_), tCrB(tCrB_) {}
|
||||
};
|
||||
|
||||
// Host side kernel arguments
|
||||
struct Arguments {
|
||||
ArrayElementA const* ptr_A{nullptr};
|
||||
StrideA dA{};
|
||||
ArrayElementB const* ptr_B{nullptr};
|
||||
StrideB dB{};
|
||||
RuntimeDataTypeA runtime_data_type_a{};
|
||||
RuntimeDataTypeB runtime_data_type_b{};
|
||||
};
|
||||
|
||||
// Device side kernel params
|
||||
struct Params {
|
||||
using ClusterLayout_VMNK = decltype(tiled_divide(make_layout(conditional_return<IsDynamicCluster>(make_shape(uint32_t(0), uint32_t(0), Int<1>{}), ClusterShape{})),
|
||||
make_tile(typename TiledMma::AtomThrID{})));
|
||||
|
||||
using TMA_A = decltype(make_tma_atom_A_sm100<TmaInternalElementA>(
|
||||
GmemTiledCopyA{},
|
||||
make_tensor(recast_ptr<TmaInternalElementA>(nullptr), repeat_like(StrideA{}, int32_t(0)), StrideA{}),
|
||||
SmemLayoutA{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
ClusterLayout_VMNK{})
|
||||
);
|
||||
|
||||
using TMA_B = decltype(make_tma_atom_B_sm100<TmaInternalElementB>(
|
||||
GmemTiledCopyB{},
|
||||
make_tensor(recast_ptr<TmaInternalElementB>(nullptr), repeat_like(StrideB{}, int32_t(0)), StrideB{}),
|
||||
SmemLayoutB{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
ClusterLayout_VMNK{})
|
||||
);
|
||||
|
||||
TMA_A tma_load_a;
|
||||
TMA_B tma_load_b;
|
||||
TMA_A tma_load_a_fallback;
|
||||
TMA_B tma_load_b_fallback;
|
||||
dim3 cluster_shape_fallback;
|
||||
RuntimeDataTypeA runtime_data_type_a;
|
||||
RuntimeDataTypeB runtime_data_type_b;
|
||||
};
|
||||
|
||||
CUTLASS_DEVICE
|
||||
CollectiveMma(Params const& params, ClusterShape cluster_shape, uint32_t block_rank_in_cluster)
|
||||
: cluster_shape_(cluster_shape)
|
||||
, block_rank_in_cluster_(block_rank_in_cluster)
|
||||
, runtime_data_type_a_(params.runtime_data_type_a)
|
||||
, runtime_data_type_b_(params.runtime_data_type_b) {
|
||||
if constexpr (IsDynamicCluster) {
|
||||
const bool is_fallback_cluster = (cute::size<0>(cluster_shape_) == params.cluster_shape_fallback.x &&
|
||||
cute::size<1>(cluster_shape_) == params.cluster_shape_fallback.y);
|
||||
observed_tma_load_a_ = is_fallback_cluster ? ¶ms.tma_load_a_fallback : ¶ms.tma_load_a;
|
||||
observed_tma_load_b_ = is_fallback_cluster ? ¶ms.tma_load_b_fallback : ¶ms.tma_load_b;
|
||||
}
|
||||
else {
|
||||
observed_tma_load_a_ = ¶ms.tma_load_a;
|
||||
observed_tma_load_b_ = ¶ms.tma_load_b;
|
||||
}
|
||||
}
|
||||
|
||||
template <class ProblemShape>
|
||||
static constexpr Params
|
||||
to_underlying_arguments(
|
||||
ProblemShape const& problem_shape,
|
||||
Arguments const& args,
|
||||
[[maybe_unused]] void* workspace,
|
||||
cutlass::KernelHardwareInfo const& hw_info = cutlass::KernelHardwareInfo{}) {
|
||||
|
||||
// Optionally append 1s until problem shape is rank-4 (MNKL), in case it is only rank-3 (MNK)
|
||||
auto problem_shape_MNKL = append<4>(problem_shape, 1);
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
|
||||
auto ptr_A = recast_ptr<TmaInternalElementA>(args.ptr_A);
|
||||
auto ptr_B = recast_ptr<TmaInternalElementB>(args.ptr_B);
|
||||
|
||||
Tensor tensor_a = make_tensor(ptr_A, make_layout(make_shape(M,K,L), args.dA));
|
||||
Tensor tensor_b = make_tensor(ptr_B, make_layout(make_shape(N,K,L), args.dB));
|
||||
|
||||
auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape);
|
||||
|
||||
// Cluster layout for TMA construction
|
||||
auto cluster_layout_vmnk = tiled_divide(make_layout(cluster_shape), make_tile(typename TiledMma::AtomThrID{}));
|
||||
auto cluster_shape_fallback = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape_fallback);
|
||||
auto cluster_layout_vmnk_fallback = tiled_divide(make_layout(cluster_shape_fallback), make_tile(typename TiledMma::AtomThrID{}));
|
||||
typename Params::TMA_A tma_load_a = make_tma_atom_A_sm100<TmaInternalElementA>(
|
||||
GmemTiledCopyA{},
|
||||
tensor_a,
|
||||
SmemLayoutA{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
cluster_layout_vmnk);
|
||||
|
||||
typename Params::TMA_B tma_load_b = make_tma_atom_B_sm100<TmaInternalElementB>(
|
||||
GmemTiledCopyB{},
|
||||
tensor_b,
|
||||
SmemLayoutB{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
cluster_layout_vmnk);
|
||||
|
||||
typename Params::TMA_A tma_load_a_fallback = make_tma_atom_A_sm100<TmaInternalElementA>(
|
||||
GmemTiledCopyA{},
|
||||
tensor_a,
|
||||
SmemLayoutA{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
cluster_layout_vmnk_fallback);
|
||||
|
||||
typename Params::TMA_B tma_load_b_fallback = make_tma_atom_B_sm100<TmaInternalElementB>(
|
||||
GmemTiledCopyB{},
|
||||
tensor_b,
|
||||
SmemLayoutB{}(_,_,_,cute::Int<0>{}),
|
||||
TileShape{},
|
||||
TiledMma{},
|
||||
cluster_layout_vmnk_fallback);
|
||||
|
||||
return {
|
||||
tma_load_a,
|
||||
tma_load_b,
|
||||
tma_load_a_fallback,
|
||||
tma_load_b_fallback,
|
||||
hw_info.cluster_shape_fallback,
|
||||
args.runtime_data_type_a,
|
||||
args.runtime_data_type_b
|
||||
};
|
||||
}
|
||||
|
||||
template <class ProblemShape>
|
||||
static bool
|
||||
can_implement(
|
||||
ProblemShape const& problem_shape,
|
||||
[[maybe_unused]] Arguments const& args) {
|
||||
auto problem_shape_MNKL = append<4>(problem_shape, 1);
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
|
||||
static constexpr bool IsF8F6F4 = detail::is_sm100_mma_f8f6f4<TiledMma, ElementA, ElementB>();
|
||||
constexpr int tma_alignment_bits_A = cutlass::detail::get_input_alignment_bits<ElementA, IsF8F6F4>();
|
||||
constexpr int tma_alignment_bits_B = cutlass::detail::get_input_alignment_bits<ElementB, IsF8F6F4>();
|
||||
constexpr int min_tma_aligned_elements_A = tma_alignment_bits_A / cute::sizeof_bits<ElementA>::value;
|
||||
|
||||
bool implementable = true;
|
||||
implementable = implementable && cutlass::detail::check_alignment<min_tma_aligned_elements_A>(cute::make_shape(M,K,L), StrideA{});
|
||||
constexpr int min_tma_aligned_elements_B = tma_alignment_bits_B / cute::sizeof_bits<ElementB>::value;
|
||||
implementable = implementable && cutlass::detail::check_alignment<min_tma_aligned_elements_B>(cute::make_shape(N,K,L), StrideB{});
|
||||
|
||||
if (!implementable) {
|
||||
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Problem Size doesn't meet the minimum alignment requirements for TMA.\n");
|
||||
}
|
||||
return implementable;
|
||||
}
|
||||
|
||||
/// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance
|
||||
CUTLASS_DEVICE void
|
||||
prefetch_tma_descriptors() {
|
||||
cute::prefetch_tma_descriptor(observed_tma_load_a_->get_tma_descriptor());
|
||||
cute::prefetch_tma_descriptor(observed_tma_load_b_->get_tma_descriptor());
|
||||
}
|
||||
|
||||
/// Construct A Single Stage's Accumulator Shape
|
||||
CUTLASS_DEVICE static
|
||||
auto
|
||||
partition_accumulator_shape() {
|
||||
auto acc_shape = partition_shape_C(TiledMma{}, take<0,2>(TileShape{})); // ((MMA_TILE_M,MMA_TILE_N),MMA_M,MMA_N)
|
||||
|
||||
return acc_shape;
|
||||
}
|
||||
|
||||
template <class TmemStorage>
|
||||
CUTLASS_DEVICE static
|
||||
auto
|
||||
slice_accumulator(TmemStorage tmem_storage, int stage) {
|
||||
return cute::make_tuple(tmem_storage.accumulators(_,_,_,stage));
|
||||
}
|
||||
|
||||
template<class EpilogueTile, bool IsOverlappingAccum = false>
|
||||
CUTLASS_DEVICE static
|
||||
auto
|
||||
init_tmem_tensors(EpilogueTile epi_tile) {
|
||||
TiledMma tiled_mma;
|
||||
auto acc_shape = partition_accumulator_shape();
|
||||
// ((MMA_TILE_M,MMA_TILE_N),MMA_M,MMA_N,ACC_PIPE) where ACC_PIPE=2 so we can double buffer our accumulators for mainloop and epilogue.
|
||||
Tensor accumulators = cutlass::detail::make_sm100_accumulator<AccumulatorPipelineStageCount, IsOverlappingAccum>(
|
||||
tiled_mma, acc_shape, EpilogueTile{});
|
||||
TmemStorage<decltype(accumulators)> tmem_storage;
|
||||
tmem_storage.accumulators = accumulators;
|
||||
return tmem_storage;
|
||||
}
|
||||
|
||||
template<class AccTensor>
|
||||
CUTLASS_DEVICE static
|
||||
void
|
||||
set_tmem_offsets(TmemStorage<AccTensor>& tmem_storage, uint32_t tmem_base_addr) {
|
||||
tmem_storage.accumulators.data() = tmem_base_addr;
|
||||
}
|
||||
|
||||
/// Set up the data needed by this collective for load.
|
||||
/// Return tuple element contain
|
||||
/// gA_mkl - The tiled tma tensor for input A
|
||||
/// gB_nkl - The tiled tma tensor for input B
|
||||
/// tAsA - partitioned smem tensor for A
|
||||
/// tBsB - partitioned smem tensor for B
|
||||
/// mcast_mask_a - tma multicast mask for A
|
||||
/// mcast_mask_b - tma multicast mask for B
|
||||
template <class ProblemShape_MNKL>
|
||||
CUTLASS_DEVICE auto
|
||||
load_init(
|
||||
ProblemShape_MNKL const& problem_shape_MNKL,
|
||||
TensorStorage& shared_tensors) const {
|
||||
using X = Underscore;
|
||||
|
||||
// Separate out problem shape for convenience
|
||||
auto [M,N,K,L] = problem_shape_MNKL;
|
||||
|
||||
// Represent the full tensors -- get these from TMA
|
||||
Tensor mA_mkl = observed_tma_load_a_->get_tma_tensor(make_shape(M,K,L));
|
||||
Tensor mB_nkl = observed_tma_load_b_->get_tma_tensor(make_shape(N,K,L));
|
||||
|
||||
// Tile the tensors and defer the slice
|
||||
Tensor gA_mkl = local_tile(mA_mkl, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M, BLK_K, m, k, l)
|
||||
Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N, BLK_K, n, k, l)
|
||||
|
||||
// Partition for this CTA
|
||||
ThrMMA cta_mma = TiledMma{}.get_slice(blockIdx.x % size(typename TiledMma::AtomThrID{}));
|
||||
|
||||
Tensor tCgA_mkl = cta_mma.partition_A(gA_mkl); // (MMA, MMA_M, MMA_K, m, k, l)
|
||||
Tensor tCgB_nkl = cta_mma.partition_B(gB_nkl); // (MMA, MMA_N, MMA_K, n, k, l)
|
||||
|
||||
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (MMA,MMA_M,MMA_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (MMA,MMA_N,MMA_K,PIPE)
|
||||
|
||||
// Define the CTA-in-cluster Layout and Coord
|
||||
Layout cta_layout_mnk = make_layout(cluster_shape_);
|
||||
Layout cta_layout_vmnk = tiled_divide(cta_layout_mnk, make_tile(typename TiledMma::AtomThrID{}));
|
||||
auto cta_coord_vmnk = cta_layout_vmnk.get_flat_coord(block_rank_in_cluster_);
|
||||
|
||||
// Project the cta_layout for tma_a along the n-modes
|
||||
auto [tAgA_mkl, tAsA] = tma_partition(*observed_tma_load_a_,
|
||||
get<2>(cta_coord_vmnk), make_layout(size<2>(cta_layout_vmnk)),
|
||||
group_modes<0,3>(sA), group_modes<0,3>(tCgA_mkl));
|
||||
|
||||
// Project the cta_layout for tma_b along the m-modes
|
||||
auto [tBgB_nkl, tBsB] = tma_partition(*observed_tma_load_b_,
|
||||
get<1>(cta_coord_vmnk), make_layout(size<1>(cta_layout_vmnk)),
|
||||
group_modes<0,3>(sB), group_modes<0,3>(tCgB_nkl));
|
||||
|
||||
// TMA Multicast Masks
|
||||
uint16_t mcast_mask_a = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk);
|
||||
uint16_t mcast_mask_b = create_tma_multicast_mask<1>(cta_layout_vmnk, cta_coord_vmnk);
|
||||
|
||||
LoadParams load_params {
|
||||
shape<3>(gA_mkl), // for scheduler
|
||||
tAgA_mkl, tBgB_nkl, tAsA, tBsB, // for input tensor values
|
||||
mcast_mask_a, mcast_mask_b // multicast masks
|
||||
};
|
||||
return load_params;
|
||||
}
|
||||
|
||||
/// Set up the data needed by this collective for mma compute.
|
||||
template <class AccTensor>
|
||||
CUTLASS_DEVICE auto
|
||||
mma_init(
|
||||
[[maybe_unused]] TmemStorage<AccTensor> tmem_tensors,
|
||||
TensorStorage& shared_tensors) const {
|
||||
Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE)
|
||||
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE)
|
||||
|
||||
// Allocate "fragments/descriptors" for A and B matrices
|
||||
Tensor tCrA = TiledMma::make_fragment_A(sA); // (MMA,MMA_M,MMA_K,PIPE)
|
||||
Tensor tCrB = TiledMma::make_fragment_B(sB); // (MMA,MMA_N,MMA_K,PIPE)
|
||||
|
||||
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sA)); // PIPE
|
||||
CUTE_STATIC_ASSERT_V(Int<DispatchPolicy::Stages>{} == size<3>(sB));
|
||||
|
||||
TiledMma tiled_mma;
|
||||
|
||||
if constexpr (IsRuntimeDataType) {
|
||||
// Update instruction descriptor according to runtime argument.
|
||||
// Applying bitmask (0b111) to help compiler deduce that the conversion and assignment are safe.
|
||||
tiled_mma.idesc_.a_format_ = uint8_t(runtime_data_type_a_) & 0b111;
|
||||
tiled_mma.idesc_.b_format_ = uint8_t(runtime_data_type_b_) & 0b111;
|
||||
}
|
||||
MmaParams<decltype(tCrA), decltype(tCrB)> mma_params {
|
||||
tiled_mma,
|
||||
tCrA, tCrB
|
||||
};
|
||||
return mma_params;
|
||||
}
|
||||
|
||||
/// Perform a collective-scoped matrix multiply-accumulate
|
||||
/// Producer Perspective
|
||||
template <
|
||||
class LoadParams,
|
||||
class TileCoordMNKL,
|
||||
class KTileIterator
|
||||
>
|
||||
CUTLASS_DEVICE auto
|
||||
load(
|
||||
MainloopPipeline mainloop_pipeline,
|
||||
MainloopPipelineState mainloop_pipe_producer_state,
|
||||
LoadParams const& load_inputs,
|
||||
TileCoordMNKL const& cta_coord_mnkl,
|
||||
KTileIterator k_tile_iter, int k_tile_count) {
|
||||
|
||||
auto [unused_k_tiles,
|
||||
tAgA_mkl, tBgB_nkl, tAsA, tBsB,
|
||||
mcast_mask_a, mcast_mask_b] = load_inputs;
|
||||
|
||||
// slice out the work coord from partitioned tensors
|
||||
Tensor tAgA = tAgA_mkl(_, get<0>(cta_coord_mnkl) / size(typename TiledMma::AtomThrID{}), _, get<3>(cta_coord_mnkl));
|
||||
Tensor tBgB = tBgB_nkl(_, get<1>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl));
|
||||
|
||||
auto barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
|
||||
|
||||
// Issue the Mainloop loads
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
while (k_tile_count > 0) {
|
||||
// LOCK mainloop_pipe_producer_state for _writing_
|
||||
mainloop_pipeline.producer_acquire(mainloop_pipe_producer_state, barrier_token);
|
||||
|
||||
using BarrierType = typename MainloopPipeline::ProducerBarrierType;
|
||||
BarrierType* tma_barrier = mainloop_pipeline.producer_get_barrier(mainloop_pipe_producer_state);
|
||||
|
||||
int write_stage = mainloop_pipe_producer_state.index();
|
||||
++mainloop_pipe_producer_state;
|
||||
barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
|
||||
|
||||
if (cute::elect_one_sync()) {
|
||||
copy(observed_tma_load_a_->with(*tma_barrier, mcast_mask_a), tAgA(_,*k_tile_iter), tAsA(_,write_stage));
|
||||
copy(observed_tma_load_b_->with(*tma_barrier, mcast_mask_b), tBgB(_,*k_tile_iter), tBsB(_,write_stage));
|
||||
}
|
||||
|
||||
--k_tile_count;
|
||||
++k_tile_iter;
|
||||
}
|
||||
|
||||
return cute::make_tuple(mainloop_pipe_producer_state, k_tile_iter);
|
||||
}
|
||||
|
||||
/// Perform a Producer Epilogue to prevent early exit of ctas in a Cluster
|
||||
CUTLASS_DEVICE void
|
||||
load_tail(MainloopPipeline mainloop_pipeline, MainloopPipelineState mainloop_pipe_producer_state) {
|
||||
// Issue the epilogue waits
|
||||
// This helps avoid early exit of ctas in Cluster
|
||||
// Waits for all stages to either be released (all
|
||||
// Consumer UNLOCKs), or if the stage was never used
|
||||
// then would just be acquired since the phase was
|
||||
// still inverted from make_producer_start_state
|
||||
mainloop_pipeline.producer_tail(mainloop_pipe_producer_state);
|
||||
}
|
||||
|
||||
/// Perform a collective-scoped matrix multiply-accumulate
|
||||
/// Consumer Perspective
|
||||
template <
|
||||
class AccumulatorPipeline,
|
||||
class FrgEngine, class FrgLayout,
|
||||
class MmaParams,
|
||||
class CtaTileCoord
|
||||
>
|
||||
CUTLASS_DEVICE auto
|
||||
mma(cute::tuple<MainloopPipeline,
|
||||
AccumulatorPipeline> pipelines,
|
||||
cute::tuple<MainloopPipelineState,
|
||||
typename AccumulatorPipeline::PipelineState> pipeline_states,
|
||||
cute::tuple<cute::Tensor<FrgEngine, FrgLayout>> const& accumulators_pair,
|
||||
MmaParams const& mma_inputs,
|
||||
CtaTileCoord cta_tile_coord,
|
||||
int k_tile_count
|
||||
) {
|
||||
static_assert(is_tmem<FrgEngine>::value, "Accumulator must be tmem resident.");
|
||||
static_assert(rank(FrgLayout{}) == 3, "Accumulator must be MMA-partitioned: (MMA, MMA_M, MMA_N)");
|
||||
auto accumulators = get<0>(accumulators_pair);
|
||||
auto [tiled_mma, tCrA, tCrB] = mma_inputs;
|
||||
|
||||
auto [mainloop_pipeline, accumulator_pipeline] = pipelines;
|
||||
auto [mainloop_pipe_consumer_state, accumulator_pipe_producer_state] = pipeline_states;
|
||||
|
||||
uint32_t skip_wait = k_tile_count <= 0;
|
||||
auto barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
|
||||
|
||||
//
|
||||
// PIPELINED MAIN LOOP
|
||||
//
|
||||
tiled_mma.accumulate_ = UMMA::ScaleOut::Zero;
|
||||
CUTLASS_PRAGMA_NO_UNROLL
|
||||
while (k_tile_count > 0) {
|
||||
// WAIT on mainloop_pipe_consumer_state until its data are available
|
||||
// (phase bit flips from mainloop_pipe_consumer_state.phase() value)
|
||||
mainloop_pipeline.consumer_wait(mainloop_pipe_consumer_state, barrier_token);
|
||||
|
||||
// Compute on k_tile
|
||||
int read_stage = mainloop_pipe_consumer_state.index();
|
||||
// Save current mainlop pipeline read state
|
||||
auto curr_mainloop_pipe_consumer_state = mainloop_pipe_consumer_state;
|
||||
|
||||
// Advance mainloop_pipe
|
||||
++mainloop_pipe_consumer_state;
|
||||
--k_tile_count;
|
||||
skip_wait = k_tile_count <= 0;
|
||||
// Peek at next iteration
|
||||
barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait);
|
||||
|
||||
// Unroll the K mode manually so we can set scale C to 1
|
||||
CUTLASS_PRAGMA_UNROLL
|
||||
for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) {
|
||||
// (V,M) x (V,N) => (V,M,N)
|
||||
cute::gemm(tiled_mma,
|
||||
tCrA(_,_,k_block,read_stage),
|
||||
tCrB(_,_,k_block,read_stage),
|
||||
accumulators);
|
||||
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
|
||||
}
|
||||
mainloop_pipeline.consumer_release(curr_mainloop_pipe_consumer_state);
|
||||
}
|
||||
|
||||
return mainloop_pipe_consumer_state;
|
||||
}
|
||||
|
||||
private:
|
||||
|
||||
typename Params::TMA_A const* observed_tma_load_a_{nullptr};
|
||||
typename Params::TMA_B const* observed_tma_load_b_{nullptr};
|
||||
|
||||
RuntimeDataTypeA runtime_data_type_a_{};
|
||||
RuntimeDataTypeB runtime_data_type_b_{};
|
||||
|
||||
ClusterShape cluster_shape_;
|
||||
uint32_t block_rank_in_cluster_;
|
||||
};
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::collective
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
@@ -388,6 +388,17 @@ public:
|
||||
[[maybe_unused]] dim3 cluster(cute::size<0>(typename GemmKernel::DispatchPolicy::ClusterShape{}),
|
||||
cute::size<1>(typename GemmKernel::DispatchPolicy::ClusterShape{}),
|
||||
cute::size<2>(typename GemmKernel::DispatchPolicy::ClusterShape{}));
|
||||
|
||||
// Dynamic cluster support
|
||||
[[maybe_unused]] dim3 fallback_cluster = dim3{0,0,0};
|
||||
if constexpr (GemmKernel::ArchTag::kMinComputeCapability == 100
|
||||
) {
|
||||
if constexpr (!cute::is_static_v<typename GemmKernel::DispatchPolicy::ClusterShape>) {
|
||||
fallback_cluster = params.hw_info.cluster_shape_fallback;
|
||||
cluster = params.hw_info.cluster_shape;
|
||||
}
|
||||
}
|
||||
|
||||
[[maybe_unused]] void* kernel_params[] = {¶ms};
|
||||
|
||||
if constexpr (kEnableCudaHostAdapter) {
|
||||
@@ -415,6 +426,7 @@ public:
|
||||
else {
|
||||
launch_result = cuda_adapter->launch(grid,
|
||||
cluster,
|
||||
fallback_cluster,
|
||||
block,
|
||||
smem_size,
|
||||
stream,
|
||||
@@ -430,8 +442,7 @@ public:
|
||||
else {
|
||||
CUTLASS_ASSERT(cuda_adapter == nullptr);
|
||||
[[maybe_unused]] void const* kernel = (void const*) device_kernel<GemmKernel>;
|
||||
static constexpr bool kClusterLaunch = GemmKernel::ArchTag::kMinComputeCapability == 90
|
||||
;
|
||||
static constexpr bool kClusterLaunch = GemmKernel::ArchTag::kMinComputeCapability == 90;
|
||||
if constexpr (kClusterLaunch) {
|
||||
if constexpr (is_static_1x1x1) {
|
||||
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
|
||||
@@ -456,6 +467,42 @@ public:
|
||||
grid, cluster, block, smem_size, stream, kernel, kernel_params, launch_with_pdl);
|
||||
}
|
||||
}
|
||||
|
||||
else {
|
||||
if constexpr (GemmKernel::ArchTag::kMinComputeCapability == 100
|
||||
) {
|
||||
if constexpr (is_static_1x1x1) {
|
||||
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::run: Launching static 1x1x1 kernel");
|
||||
#endif
|
||||
launch_result = cutlass::kernel_launch<GemmKernel>(grid, block, smem_size, stream, params, launch_with_pdl);
|
||||
if (launch_result != Status::kSuccess) {
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::run: cutlass::kernel_launch reports failure");
|
||||
}
|
||||
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
|
||||
else {
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::run: cutlass::kernel_launch reports success");
|
||||
}
|
||||
#endif
|
||||
}
|
||||
else {
|
||||
#if (CUTLASS_DEBUG_TRACE_LEVEL > 1)
|
||||
CUTLASS_TRACE_HOST("GemmUniversal::run: Launching kernel with fall-back cluster");
|
||||
#endif
|
||||
launch_result = ClusterLauncher::launch_with_fallback_cluster(
|
||||
grid,
|
||||
cluster,
|
||||
fallback_cluster,
|
||||
block,
|
||||
smem_size,
|
||||
stream,
|
||||
kernel,
|
||||
kernel_params,
|
||||
launch_with_pdl);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
}
|
||||
else {
|
||||
|
||||
@@ -35,6 +35,8 @@
|
||||
|
||||
#include "cute/layout.hpp"
|
||||
#include "cute/numeric/integral_constant.hpp" // cute::false_type
|
||||
#include "cute/arch/copy_sm100.hpp"
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::detail {
|
||||
@@ -346,6 +348,174 @@ struct MainloopSm90TmaGmmaWarpSpecializedSparseFP8
|
||||
};
|
||||
|
||||
|
||||
template<
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_
|
||||
>
|
||||
struct KernelTmaWarpSpecializedSm100 final {
|
||||
static constexpr int SchedulerPipelineStageCount = SchedulerPipelineStageCount_;
|
||||
static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_;
|
||||
};
|
||||
|
||||
// Gemm with block scaling factors
|
||||
template<
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_
|
||||
>
|
||||
struct KernelTmaWarpSpecializedBlockScaledSm100 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_,
|
||||
int AccumulatorPipelineStageCount_
|
||||
>
|
||||
struct KernelPtrArrayTmaWarpSpecializedSm100 final {
|
||||
static constexpr int SchedulerPipelineStageCount = SchedulerPipelineStageCount_;
|
||||
static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_;
|
||||
};
|
||||
|
||||
// Ptr-Array Block Scaled GEMM
|
||||
template<
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_
|
||||
>
|
||||
struct KernelPtrArrayTmaWarpSpecializedBlockScaledSm100 final {
|
||||
static constexpr int SchedulerPipelineStageCount = SchedulerPipelineStageCount_;
|
||||
static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_;
|
||||
};
|
||||
|
||||
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
//
|
||||
// Collective Builder Tag Property
|
||||
//
|
||||
|
||||
struct KernelSchedule1Sm {};
|
||||
struct KernelSchedule2Sm {};
|
||||
struct KernelScheduleSm100 {};
|
||||
struct KernelScheduleSm100DenseGemm : KernelScheduleSm100 {};
|
||||
|
||||
struct KernelScheduleBlockScaledGemmSm100 : KernelScheduleSm100 {};
|
||||
struct KernelScheduleMxNvf4Sm100 : KernelScheduleBlockScaledGemmSm100 {};
|
||||
struct KernelScheduleMxf8f6f4Sm100 : KernelScheduleBlockScaledGemmSm100 {};
|
||||
|
||||
struct KernelScheduleSm100PtrArrayDenseGemm : KernelScheduleSm100DenseGemm {};
|
||||
struct KernelSchedulePtrArrayBlockScaledGemmSm100 : KernelScheduleBlockScaledGemmSm100 {};
|
||||
struct KernelSchedulePtrArrayMxNvf4Sm100 : KernelSchedulePtrArrayBlockScaledGemmSm100 {};
|
||||
struct KernelSchedulePtrArrayMxf8f6f4Sm100 : KernelSchedulePtrArrayBlockScaledGemmSm100 {};
|
||||
|
||||
|
||||
//
|
||||
// Collective Builder Tag
|
||||
// Only used in CollectiveBuilder
|
||||
//
|
||||
|
||||
// Dense GEMM: Specialize for 1SM vs 2SM
|
||||
struct KernelTmaWarpSpecialized1SmSm100 final : KernelSchedule1Sm, KernelScheduleSm100DenseGemm {};
|
||||
struct KernelTmaWarpSpecialized2SmSm100 final : KernelSchedule2Sm, KernelScheduleSm100DenseGemm {};
|
||||
|
||||
|
||||
|
||||
// Block Scaled Dense GEMM: Specialize for instruction type, scale factor vector size, and 1SM vs. 2SM
|
||||
struct KernelTmaWarpSpecialized1SmBlockScaledSm100 final : KernelSchedule1Sm, KernelScheduleBlockScaledGemmSm100 { };
|
||||
struct KernelTmaWarpSpecialized2SmBlockScaledSm100 final : KernelSchedule2Sm, KernelScheduleBlockScaledGemmSm100 { };
|
||||
struct KernelTmaWarpSpecialized1SmNvf4Sm100 final : KernelSchedule1Sm, KernelScheduleMxNvf4Sm100 { };
|
||||
struct KernelTmaWarpSpecialized2SmNvf4Sm100 final : KernelSchedule2Sm, KernelScheduleMxNvf4Sm100 { };
|
||||
struct KernelTmaWarpSpecialized1SmMxf4Sm100 final : KernelSchedule1Sm, KernelScheduleMxNvf4Sm100 { };
|
||||
struct KernelTmaWarpSpecialized2SmMxf4Sm100 final : KernelSchedule2Sm, KernelScheduleMxNvf4Sm100 { };
|
||||
struct KernelTmaWarpSpecialized1SmMxf8f6f4Sm100 final : KernelSchedule1Sm, KernelScheduleMxf8f6f4Sm100 { };
|
||||
struct KernelTmaWarpSpecialized2SmMxf8f6f4Sm100 final : KernelSchedule2Sm, KernelScheduleMxf8f6f4Sm100 { };
|
||||
|
||||
|
||||
// Ptr-Array Dense GEMM: Specialize for 1SM vs 2SM
|
||||
struct KernelPtrArrayTmaWarpSpecialized1SmSm100 final : KernelSchedule1Sm, KernelScheduleSm100PtrArrayDenseGemm {};
|
||||
struct KernelPtrArrayTmaWarpSpecialized2SmSm100 final : KernelSchedule2Sm, KernelScheduleSm100PtrArrayDenseGemm {};
|
||||
|
||||
|
||||
// Ptr-Array Block Scaled Dense GEMM: Specialize for instruction type, scale factor vector size, and 1SM vs. 2SM
|
||||
struct KernelPtrArrayTmaWarpSpecialized1SmBlockScaledSm100 final : KernelSchedule1Sm, KernelSchedulePtrArrayBlockScaledGemmSm100 { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized2SmBlockScaledSm100 final : KernelSchedule2Sm, KernelSchedulePtrArrayBlockScaledGemmSm100 { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized1SmNvf4Sm100 final : KernelSchedule1Sm, KernelSchedulePtrArrayMxNvf4Sm100 { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized2SmNvf4Sm100 final : KernelSchedule2Sm, KernelSchedulePtrArrayMxNvf4Sm100 { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized1SmMxf4Sm100 final : KernelSchedule1Sm, KernelSchedulePtrArrayMxNvf4Sm100 { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized2SmMxf4Sm100 final : KernelSchedule2Sm, KernelSchedulePtrArrayMxNvf4Sm100 { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized1SmMxf8f6f4Sm100 final : KernelSchedule1Sm, KernelSchedulePtrArrayMxf8f6f4Sm100 { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized2SmMxf8f6f4Sm100 final : KernelSchedule2Sm, KernelSchedulePtrArrayMxf8f6f4Sm100 { };
|
||||
|
||||
|
||||
// n-buffer in smem, pipelined with Blackwell UMMA and TMA, Warp specialized dynamic schedule
|
||||
template<
|
||||
int Stages_,
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>
|
||||
>
|
||||
struct MainloopSm100TmaUmmaWarpSpecialized {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using ArchTag = arch::Sm100;
|
||||
using Schedule = KernelTmaWarpSpecializedSm100<SchedulerPipelineStageCount_, AccumulatorPipelineStageCount_>;
|
||||
constexpr static bool IsOverlappingAccum = false;
|
||||
};
|
||||
|
||||
|
||||
|
||||
// n-buffer in smem, pipelined with Blackwell UMMA and TMA, Warp specialized dynamic schedule
|
||||
template<
|
||||
int Stages_,
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>
|
||||
>
|
||||
struct MainloopSm100TmaUmmaWarpSpecializedBlockScaled {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using ArchTag = arch::Sm100;
|
||||
constexpr static bool IsOverlappingAccum = AccumulatorPipelineStageCount_ == 1;
|
||||
using Schedule = KernelTmaWarpSpecializedBlockScaledSm100<SchedulerPipelineStageCount_, AccumulatorPipelineStageCount_>;
|
||||
};
|
||||
|
||||
|
||||
|
||||
// n-buffer in smem, pipelined with Blackwell UMMA and TMA, Warp specialized dynamic schedule
|
||||
template<
|
||||
int Stages_,
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>
|
||||
>
|
||||
struct MainloopSm100ArrayTmaUmmaWarpSpecialized {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using ArchTag = arch::Sm100;
|
||||
constexpr static bool IsOverlappingAccum = false;
|
||||
using Schedule = KernelPtrArrayTmaWarpSpecializedSm100<SchedulerPipelineStageCount_, AccumulatorPipelineStageCount_>;
|
||||
};
|
||||
|
||||
// n-buffer in smem, pipelined with Blackwell UMMA and TMA, Warp specialized dynamic schedule
|
||||
template<
|
||||
int Stages_,
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>
|
||||
>
|
||||
struct MainloopSm100ArrayTmaUmmaWarpSpecializedBlockScaled {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using ArchTag = arch::Sm100;
|
||||
constexpr static bool IsOverlappingAccum = AccumulatorPipelineStageCount_ == 1;
|
||||
using Schedule = KernelPtrArrayTmaWarpSpecializedBlockScaledSm100<SchedulerPipelineStageCount_, AccumulatorPipelineStageCount_>;
|
||||
};
|
||||
|
||||
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm
|
||||
|
||||
@@ -63,4 +63,6 @@ struct IsCutlass3ArrayKernel<ProblemShape, cute::void_t<typename ProblemShape::U
|
||||
#include "cutlass/gemm/kernel/sm90_gemm_tma_warpspecialized_cooperative.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_gemm_array_tma_warpspecialized_pingpong.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_gemm_array_tma_warpspecialized_cooperative.hpp"
|
||||
#include "cutlass/gemm/kernel/sm100_gemm_tma_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/kernel/sm100_gemm_array_tma_warpspecialized.hpp"
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
+723
@@ -0,0 +1,723 @@
|
||||
/***************************************************************************************************
|
||||
* 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 "cute/int_tuple.hpp"
|
||||
|
||||
#include "cutlass/arch/config.h"
|
||||
#include "cutlass/arch/barrier.h"
|
||||
#include "cutlass/detail/cluster.hpp"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/gemm_coord.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_tile_scheduler.hpp"
|
||||
#include "cutlass/gemm/kernel/tile_scheduler_params.h"
|
||||
#include "cutlass/conv/convnd_problem_shape.hpp"
|
||||
#include "cutlass/conv/detail.hpp"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::kernel::detail {
|
||||
|
||||
//////////////////// Blackwell Scheduler /////////////////////////
|
||||
|
||||
template<
|
||||
class ClusterShape_,
|
||||
uint32_t Stages_
|
||||
>
|
||||
class PersistentTileSchedulerSm100 {
|
||||
|
||||
private:
|
||||
|
||||
using UnderlyingTileScheduler = PersistentTileSchedulerSm90;
|
||||
|
||||
public:
|
||||
|
||||
using ClusterShape = ClusterShape_;
|
||||
using RasterOrder = UnderlyingTileScheduler::RasterOrder;
|
||||
using RasterOrderOptions = UnderlyingTileScheduler::RasterOrderOptions;
|
||||
static constexpr bool IsDynamicPersistent = true;
|
||||
|
||||
static constexpr uint32_t Stages = Stages_;
|
||||
|
||||
// CLC response is an opaque 16B value
|
||||
struct CLCResponse { uint32_t data[4]; };
|
||||
|
||||
using WorkTileInfo = typename PersistentTileSchedulerSm90::WorkTileInfo;
|
||||
|
||||
using Params = PersistentTileSchedulerSm100Params;
|
||||
|
||||
using Pipeline = PipelineCLCFetchAsync<Stages, ClusterShape>;
|
||||
using PipelineStorage = typename Pipeline::SharedStorage;
|
||||
|
||||
using ThrottlePipeline = PipelineAsync<Stages>;
|
||||
using ThrottlePipelineStorage = typename ThrottlePipeline::SharedStorage;
|
||||
|
||||
class SharedStorage {
|
||||
public:
|
||||
|
||||
CUTLASS_DEVICE PipelineStorage& pipeline() { return pipeline_; }
|
||||
CUTLASS_DEVICE ThrottlePipelineStorage& throttle_pipeline() { return throttle_pipeline_; }
|
||||
CUTLASS_DEVICE CLCResponse* data() { return data_; }
|
||||
|
||||
private:
|
||||
alignas(16) PipelineStorage pipeline_;
|
||||
alignas(16) ThrottlePipelineStorage throttle_pipeline_;
|
||||
alignas(16) CLCResponse data_[Stages];
|
||||
};
|
||||
|
||||
struct Arguments {
|
||||
|
||||
Arguments() = default;
|
||||
Arguments(Arguments const&) = default;
|
||||
Arguments(Arguments&&) = default;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments&
|
||||
operator=(Arguments const&) {
|
||||
return *this;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
Arguments&
|
||||
operator=(Arguments &&) {
|
||||
return *this;
|
||||
}
|
||||
|
||||
int max_swizzle_size = 1;
|
||||
RasterOrderOptions raster_order = RasterOrderOptions::Heuristic;
|
||||
};
|
||||
|
||||
//
|
||||
// Static Host Methods
|
||||
//
|
||||
|
||||
template <class ProblemShapeMNKL, class TileShape, class ClusterShape>
|
||||
static Params
|
||||
to_underlying_arguments(
|
||||
ProblemShapeMNKL problem_shape_mnkl,
|
||||
TileShape tile_shape,
|
||||
[[maybe_unused]] ClusterShape cluster_shape,
|
||||
[[maybe_unused]] KernelHardwareInfo const& hw_info,
|
||||
[[maybe_unused]] Arguments const& args,
|
||||
[[maybe_unused]] void* workspace = nullptr,
|
||||
[[maybe_unused]] uint32_t NumEpilogueSubTiles = 1,
|
||||
[[maybe_unused]] uint32_t ktile_start_alignment_count = 1u
|
||||
) {
|
||||
|
||||
auto cs = cutlass::detail::select_cluster_shape(ClusterShape_{}, hw_info.cluster_shape);
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape, cs);
|
||||
|
||||
Params params;
|
||||
params.initialize(
|
||||
problem_blocks,
|
||||
to_gemm_coord(cs),
|
||||
hw_info,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order
|
||||
);
|
||||
return params;
|
||||
}
|
||||
|
||||
template <class ProblemShapeMNKL, class TileShape, class AtomThrShape, class ClusterShape>
|
||||
static Params
|
||||
to_underlying_arguments(
|
||||
ProblemShapeMNKL problem_shape_mnkl,
|
||||
TileShape tile_shape_mnk,
|
||||
AtomThrShape atom_thr_shape_mnk,
|
||||
ClusterShape cluster_shape_mnk,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr
|
||||
) {
|
||||
|
||||
auto selected_cluster_shape = cutlass::detail::select_cluster_shape(cluster_shape_mnk, hw_info.cluster_shape);
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape_mnk,
|
||||
atom_thr_shape_mnk, selected_cluster_shape);
|
||||
|
||||
Params params;
|
||||
params.initialize(
|
||||
problem_blocks,
|
||||
to_gemm_coord(selected_cluster_shape),
|
||||
hw_info,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order
|
||||
);
|
||||
return params;
|
||||
}
|
||||
|
||||
// Conv Specialization
|
||||
template <conv::Operator ConvOp, int NumSpatialDims, class TileShape, class AtomThrShape, class ClusterShape>
|
||||
static Params
|
||||
to_underlying_arguments(
|
||||
cutlass::conv::ConvProblemShape<ConvOp, NumSpatialDims> problem_shape,
|
||||
TileShape tile_shape_mnk,
|
||||
AtomThrShape atom_thr_shape_mnk,
|
||||
ClusterShape cluster_shape_mnk,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr
|
||||
) {
|
||||
|
||||
auto problem_shape_mnkl = [&] () {
|
||||
// Infer im2col linearization from ConvOp and TileShape
|
||||
constexpr bool is_linearized_M = (ConvOp == conv::Operator::kFprop || ConvOp == conv::Operator::kDgrad)
|
||||
&& depth<0>(TileShape{}) == _0{};
|
||||
constexpr bool is_linearized_K = ConvOp == conv::Operator::kWgrad && depth<2>(TileShape{}) == _0{};
|
||||
|
||||
if constexpr (is_linearized_M || is_linearized_K) {
|
||||
// transformation + im2col linearization
|
||||
return cutlass::conv::detail::get_linearized_problem_shape_MNKL(problem_shape);
|
||||
}
|
||||
else {
|
||||
// transformation
|
||||
return cutlass::conv::detail::get_transformed_problem_shape_MNKL(problem_shape);
|
||||
}
|
||||
}();
|
||||
|
||||
return to_underlying_arguments(
|
||||
problem_shape_mnkl,
|
||||
tile_shape_mnk,
|
||||
atom_thr_shape_mnk,
|
||||
cluster_shape_mnk,
|
||||
hw_info,
|
||||
args,
|
||||
workspace
|
||||
);
|
||||
}
|
||||
|
||||
// Given the inputs, computes the physical grid we should launch.
|
||||
template<class ProblemShapeMNKL, class BlockShape, class ClusterShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_grid_shape(
|
||||
Params const& params,
|
||||
ProblemShapeMNKL problem_shape_mnk,
|
||||
BlockShape cta_shape,
|
||||
ClusterShape cluster_shape,
|
||||
KernelHardwareInfo hw_info,
|
||||
[[maybe_unused]] Arguments arguments) {
|
||||
auto problem_shape_MNKL = append<4>(problem_shape_mnk, Int<1>{});
|
||||
auto grid = get_tiled_cta_shape_mnl(problem_shape_MNKL, cta_shape, cluster_shape);
|
||||
return possibly_transpose_grid(params.raster_order_, params.divmod_cluster_shape_m_, params.divmod_cluster_shape_n_, grid);
|
||||
}
|
||||
|
||||
// Given the inputs, computes the physical grid we should launch.
|
||||
template<class ProblemShapeMNKL, class TileShape, class AtomThrShape, class ClusterShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_grid_shape(
|
||||
Params const& params,
|
||||
ProblemShapeMNKL problem_shape_mnkl,
|
||||
TileShape tile_shape_mnk,
|
||||
AtomThrShape atom_thr_shape_mnk,
|
||||
ClusterShape cluster_shape_mnk,
|
||||
KernelHardwareInfo hw_info) {
|
||||
auto grid = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape_mnk, atom_thr_shape_mnk, cluster_shape_mnk);
|
||||
return possibly_transpose_grid(params.raster_order_, params.divmod_cluster_shape_m_, params.divmod_cluster_shape_n_, grid);
|
||||
}
|
||||
|
||||
// Possibly transpose the grid depending on rasterization order.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
possibly_transpose_grid(RasterOrder raster_order, FastDivmod divmod_cluster_shape_m, FastDivmod divmod_cluster_shape_n, dim3 grid) {
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
// Swap grid.x and grid.y for AlongN rasterization order, since the CLC scheduler
|
||||
// will schedule in AlongM order by default.
|
||||
//
|
||||
// Each grid dimension must also be a multiple of the corresponding cluster dimension,
|
||||
// so we convert the untransposed x into the number of clusters along the M mode,
|
||||
// and multiply this by cluster.n (and vice-versa for y).
|
||||
auto tmp = grid.x;
|
||||
grid.x = divmod_cluster_shape_n.divide(grid.y) * divmod_cluster_shape_m;
|
||||
grid.y = divmod_cluster_shape_m.divide(tmp) * divmod_cluster_shape_n;
|
||||
}
|
||||
return grid;
|
||||
}
|
||||
|
||||
template <class ProblemShape, class ElementAccumulator>
|
||||
static size_t
|
||||
get_workspace_size(
|
||||
Arguments const& args,
|
||||
ProblemShape problem_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
[[maybe_unused]] uint32_t reduction_warp_groups,
|
||||
[[maybe_unused]] const uint32_t epilogue_subtile = 1,
|
||||
[[maybe_unused]] uint32_t num_accumulator_mtxs = 1) {
|
||||
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape, 1);
|
||||
|
||||
auto cs = cutlass::detail::select_cluster_shape(ClusterShape_{}, hw_info.cluster_shape);
|
||||
|
||||
return Params::get_workspace_size(
|
||||
to_gemm_coord(problem_shape_mnkl),
|
||||
GemmCoord(1, 1, 1), // Tile shape. Unused.
|
||||
to_gemm_coord(cs),
|
||||
hw_info,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order
|
||||
);
|
||||
}
|
||||
|
||||
template <class ElementAccumulator, class ProblemShape, class TileShapeMNK, class AtomThrShape, class ClusterShape>
|
||||
static size_t
|
||||
get_workspace_size(Arguments const& args, ProblemShape problem_shape, TileShapeMNK, AtomThrShape, ClusterShape, KernelHardwareInfo const& hw_info,
|
||||
uint32_t reduction_warp_groups, uint32_t num_accumulator_mtxs = 1) {
|
||||
return get_workspace_size<ProblemShape, ElementAccumulator>(args, problem_shape, hw_info, reduction_warp_groups, num_accumulator_mtxs);
|
||||
}
|
||||
|
||||
template <class ProblemShape, class ElementAccumulator>
|
||||
static cutlass::Status
|
||||
initialize_workspace(
|
||||
Arguments const& args,
|
||||
void* workspace,
|
||||
cudaStream_t stream,
|
||||
ProblemShape const& problem_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
uint32_t, // reduction_warp_groups
|
||||
uint32_t = 1, // epilogue_subtile
|
||||
uint32_t = 1, // num_accumulator_mtxs
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape, 1);
|
||||
|
||||
auto cs = cutlass::detail::select_cluster_shape(ClusterShape_{}, hw_info.cluster_shape);
|
||||
|
||||
return Params::initialize_workspace(
|
||||
workspace,
|
||||
stream,
|
||||
to_gemm_coord(problem_shape_mnkl),
|
||||
GemmCoord(1, 1, 1), // Tile shape. Unused.
|
||||
to_gemm_coord(cs),
|
||||
hw_info,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
cuda_adapter
|
||||
);
|
||||
}
|
||||
|
||||
template <class ElementAccumulator, class ProblemShape, class TileShapeMNK, class AtomThrShape>
|
||||
static cutlass::Status
|
||||
initialize_workspace(
|
||||
Arguments const& args,
|
||||
void* workspace,
|
||||
cudaStream_t stream,
|
||||
ProblemShape const& problem_shape,
|
||||
TileShapeMNK,
|
||||
AtomThrShape,
|
||||
ClusterShape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
uint32_t reduction_warp_groups,
|
||||
uint32_t num_accumulator_mtxs = 1,
|
||||
CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
|
||||
return initialize_workspace<ProblemShape, ElementAccumulator>(
|
||||
args,
|
||||
workspace,
|
||||
stream,
|
||||
problem_shape,
|
||||
hw_info,
|
||||
reduction_warp_groups,
|
||||
1, // epilogue_subtile
|
||||
num_accumulator_mtxs,
|
||||
cuda_adapter
|
||||
);
|
||||
}
|
||||
|
||||
static bool
|
||||
can_implement(Arguments const& args) {
|
||||
return true;
|
||||
}
|
||||
|
||||
//
|
||||
// Constructors
|
||||
//
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileSchedulerSm100(Params const& params)
|
||||
: scheduler_params(params) {}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileSchedulerSm100(CLCResponse* clc_response_ptr, Params const& params, dim3 block_id_in_cluster)
|
||||
: clc_response_ptr_(clc_response_ptr), scheduler_params(params), block_id_in_cluster_(block_id_in_cluster) {}
|
||||
|
||||
template <class ProblemShapeMNKL, class TileShape>
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileSchedulerSm100(CLCResponse* clc_response_ptr, Params const& params, ProblemShapeMNKL problem_shape_mnkl, TileShape tile_shape, dim3 block_id_in_cluster)
|
||||
: PersistentTileSchedulerSm100(clc_response_ptr, params, block_id_in_cluster) {}
|
||||
//
|
||||
// Data Members
|
||||
//
|
||||
CLCResponse *clc_response_ptr_ = nullptr;
|
||||
Params const& scheduler_params;
|
||||
dim3 block_id_in_cluster_;
|
||||
|
||||
//
|
||||
// Work Tile API
|
||||
//
|
||||
|
||||
// Returns the initial work tile info that will be computed over
|
||||
template <class ClusterShape>
|
||||
CUTLASS_DEVICE
|
||||
static WorkTileInfo
|
||||
initial_work_tile_info(ClusterShape cluster_shape, Params const& params) {
|
||||
WorkTileInfo work_tile{
|
||||
static_cast<int32_t>((blockIdx.x / cute::size<0>(cluster_shape)) * cute::size<0>(cluster_shape)),
|
||||
static_cast<int32_t>((blockIdx.y / cute::size<1>(cluster_shape)) * cute::size<1>(cluster_shape)),
|
||||
static_cast<int32_t>((blockIdx.z / cute::size<2>(cluster_shape)) * cute::size<2>(cluster_shape)),
|
||||
true
|
||||
};
|
||||
|
||||
possibly_transpose_work_tile(work_tile, params);
|
||||
return work_tile;
|
||||
}
|
||||
|
||||
// Returns the initial work tile info that will be computed over
|
||||
template <class ClusterShape>
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
initial_work_tile_info(ClusterShape cluster_shape) {
|
||||
return initial_work_tile_info(cluster_shape, scheduler_params);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
auto
|
||||
work_tile_to_cta_coord(WorkTileInfo work_tile_info) {
|
||||
// Get every cta coord in three dimensions of the cluster
|
||||
auto [cta_m_in_cluster, cta_n_in_cluster, cta_l_in_cluster] = block_id_in_cluster_;
|
||||
return make_coord(
|
||||
work_tile_info.M_idx + static_cast<int32_t>(cta_m_in_cluster),
|
||||
work_tile_info.N_idx + static_cast<int32_t>(cta_n_in_cluster),
|
||||
_,
|
||||
work_tile_info.L_idx + static_cast<int32_t>(cta_l_in_cluster)
|
||||
);
|
||||
}
|
||||
|
||||
// Convert CTA-level work tile info to cluster-level tile coord
|
||||
CUTLASS_DEVICE
|
||||
auto
|
||||
work_tile_to_cluster_coord_mnkl(WorkTileInfo work_tile_info) const {
|
||||
// TileScheduler works at CTA-level, kernel works at cluster-level
|
||||
int m_coord = idx2crd(scheduler_params.divmod_cluster_shape_m_.divide(work_tile_info.M_idx),
|
||||
scheduler_params.problem_tiles_m_);
|
||||
int n_coord = idx2crd(scheduler_params.divmod_cluster_shape_n_.divide(work_tile_info.N_idx),
|
||||
scheduler_params.problem_tiles_n_);
|
||||
int l_coord = idx2crd(work_tile_info.L_idx,
|
||||
scheduler_params.problem_tiles_l_);
|
||||
return make_coord(m_coord, n_coord, _, l_coord);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static void
|
||||
issue_clc_query(PipelineState<Stages> state, uint32_t mbarrier_addr, CLCResponse* clc_response_ptr) {
|
||||
#if defined(CUTLASS_ARCH_CLC_ENABLED)
|
||||
uint32_t result_addr = cute::cast_smem_ptr_to_uint(reinterpret_cast<const void*>(
|
||||
&clc_response_ptr[state.index()]));
|
||||
asm volatile(
|
||||
"{\n\t"
|
||||
"clusterlaunchcontrol.try_cancel.async.shared::cta.mbarrier::complete_tx::bytes.multicast::cluster::all.b128 [%0], [%1];\n\t"
|
||||
"}\n"
|
||||
:
|
||||
: "r"(result_addr), "r"(mbarrier_addr));
|
||||
#else
|
||||
CUTLASS_NOT_IMPLEMENTED();
|
||||
#endif
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static WorkTileInfo
|
||||
work_tile_info_from_clc_response(uint32_t result_addr) {
|
||||
WorkTileInfo work_tile_info;
|
||||
uint32_t valid = 0;
|
||||
|
||||
#if defined(CUTLASS_ARCH_CLC_ENABLED)
|
||||
asm volatile(
|
||||
"{\n"
|
||||
".reg .pred p1;\n\t"
|
||||
".reg .b128 clc_result;\n\t"
|
||||
"ld.shared.b128 clc_result, [%4];\n\t"
|
||||
"clusterlaunchcontrol.query_cancel.is_canceled.pred.b128 p1, clc_result;\n\t"
|
||||
"selp.u32 %3, 1, 0, p1;\n\t"
|
||||
"@p1 clusterlaunchcontrol.query_cancel.get_first_ctaid.v4.b32.b128 {%0, %1, %2, _}, clc_result;\n\t"
|
||||
"}\n"
|
||||
: "=r"(work_tile_info.M_idx), "=r"(work_tile_info.N_idx), "=r"(work_tile_info.L_idx), "=r"(valid)
|
||||
: "r"(result_addr)
|
||||
: "memory"
|
||||
);
|
||||
|
||||
cutlass::arch::fence_view_async_shared();
|
||||
#else
|
||||
CUTLASS_NOT_IMPLEMENTED();
|
||||
#endif
|
||||
work_tile_info.is_valid_tile = (valid == 1);
|
||||
return work_tile_info;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
PipelineState<Stages>
|
||||
advance_to_next_work(Pipeline& clc_pipeline, PipelineState<Stages> clc_pipe_producer_state) const {
|
||||
uint32_t mbarrier_addr = clc_pipeline.producer_get_barrier(clc_pipe_producer_state);
|
||||
// Wait for clcID buffer to become empty with a flipped phase
|
||||
clc_pipeline.producer_acquire(clc_pipe_producer_state);
|
||||
|
||||
if (cute::elect_one_sync()) {
|
||||
issue_clc_query(clc_pipe_producer_state, mbarrier_addr, clc_response_ptr_);
|
||||
}
|
||||
|
||||
++clc_pipe_producer_state;
|
||||
return clc_pipe_producer_state;
|
||||
}
|
||||
|
||||
// Kernel helper function to get next work tile
|
||||
template <class TileSchedulerPipeline, class TileSchedulerPipelineState>
|
||||
CUTLASS_DEVICE
|
||||
auto
|
||||
fetch_next_work(
|
||||
WorkTileInfo work_tile_info,
|
||||
TileSchedulerPipeline& scheduler_pipeline,
|
||||
TileSchedulerPipelineState scheduler_pipe_consumer_state) {
|
||||
|
||||
scheduler_pipeline.consumer_wait(scheduler_pipe_consumer_state);
|
||||
auto new_work_tile_info = get_current_work(scheduler_pipe_consumer_state);
|
||||
scheduler_pipeline.consumer_release(scheduler_pipe_consumer_state);
|
||||
|
||||
// Return true to indicate that the tile scheduler pipeline state should be advanced
|
||||
return cute::make_tuple(new_work_tile_info, true);
|
||||
}
|
||||
|
||||
//
|
||||
// K Tile API
|
||||
//
|
||||
// Permute K iteration loading order from [C, S, R, T] to [S, R, T, C] for better L2 locality
|
||||
template <class ProblemShapeMNKL, class TileShape, class Shape>
|
||||
CUTLASS_DEVICE
|
||||
auto
|
||||
get_k_tile_iterator(WorkTileInfo const& work_tile_info, ProblemShapeMNKL problem_shape_MNKL, TileShape tile_shape, Shape) {
|
||||
constexpr int32_t rank_t = cute::rank<2>(ProblemShapeMNKL{});
|
||||
auto k_tiles = cute::ceil_div(cute::get<2>(problem_shape_MNKL), cute::get<2>(tile_shape));
|
||||
if constexpr (rank_t == 4) {
|
||||
return cute::make_coord_iterator<cute::Step<_3, _0, _1, _2>>(k_tiles);
|
||||
}
|
||||
else if constexpr (rank_t == 3) {
|
||||
return cute::make_coord_iterator<cute::Step<_2, _0, _1>>(k_tiles);
|
||||
}
|
||||
else if constexpr (rank_t == 2) {
|
||||
return cute::make_coord_iterator<cute::Step<_1, _0>>(k_tiles);
|
||||
}
|
||||
else {
|
||||
return cute::make_coord_iterator(k_tiles);
|
||||
}
|
||||
}
|
||||
|
||||
template <class ProblemShape, class TileShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int
|
||||
get_work_k_tile_count(WorkTileInfo const& work_tile_info, ProblemShape problem_shape, TileShape tile_shape) {
|
||||
// All work units returned by this scheduler cover the entire K iteration
|
||||
// space of the output tile assigned to the work unit.
|
||||
return cute::size(cute::ceil_div(cute::get<2>(problem_shape), cute::get<2>(tile_shape)));
|
||||
}
|
||||
|
||||
// Compatible with sm90 kernel layers
|
||||
CUTLASS_HOST_DEVICE
|
||||
static uint32_t
|
||||
get_work_k_tile_start(WorkTileInfo const&) {
|
||||
// All work units returned by this scheduler start from K tile 0
|
||||
return 0u;
|
||||
}
|
||||
|
||||
// Returns whether the block assigned this work should compute the epilogue for the corresponding
|
||||
// output tile. For the basic tile scheduler, this is always true.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
compute_epilogue(WorkTileInfo const&, Params const&) {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
compute_epilogue(WorkTileInfo const&) {
|
||||
return true;
|
||||
}
|
||||
|
||||
// Returns whether fixup is needed for `work_tile_info`. None of the work units returned by
|
||||
// this scheduler require fixup, since none of the work units partition the reduction extent.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
requires_fixup(Params const& params, WorkTileInfo const work_tile_info) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Performs the reduction across splits for a given output tile. No fixup is required for
|
||||
// work units returned by this scheduler.
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
fixup(WorkTileInfo const&, FrgTensorC&, uint32_t, uint32_t, uint32_t = 1) const { }
|
||||
|
||||
template <
|
||||
bool IsComplex,
|
||||
class TiledMma,
|
||||
class AccEngine,
|
||||
class AccLayout,
|
||||
class AccumulatorPipeline,
|
||||
class AccumulatorPipelineState,
|
||||
class CopyOpT2R
|
||||
>
|
||||
CUTLASS_DEVICE
|
||||
AccumulatorPipelineState
|
||||
fixup(
|
||||
TiledMma const& ,
|
||||
WorkTileInfo const&,
|
||||
cute::Tensor<AccEngine, AccLayout>&,
|
||||
AccumulatorPipeline,
|
||||
AccumulatorPipelineState acc_pipe_consumer_state,
|
||||
CopyOpT2R) const {
|
||||
return acc_pipe_consumer_state;
|
||||
}
|
||||
|
||||
// Returns whether the current WorkTileInfo passed in should continue to be used. Since
|
||||
// this scheduler only schedules work in units of single, full output tiles, the WorkTileInfo
|
||||
// passed in should not be used after having been processed.
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
continue_current_work(WorkTileInfo&) {
|
||||
return false;
|
||||
}
|
||||
|
||||
//
|
||||
// Implementation Helpers
|
||||
//
|
||||
// Given the inputs, computes the total number of output blocks this problem will compute over
|
||||
// Note that this is only the logical size of our grid, not the physical grid we will actually launch.
|
||||
template<class ProblemShapeMNKL, class BlockShape, class ClusterShape>
|
||||
CUTLASS_HOST_DEVICE static dim3
|
||||
get_tiled_cta_shape_mnl(ProblemShapeMNKL problem_shape_mnkl, BlockShape blk_shape, ClusterShape cluster_shape) {
|
||||
auto grid_shape = shape(ceil_div(problem_shape_mnkl, blk_shape));
|
||||
auto grid_shape_up = round_up(product_each(grid_shape), cluster_shape); // Assumes ClusterShape is flat
|
||||
return dim3(size<0>(grid_shape_up), // M
|
||||
size<1>(grid_shape_up), // N
|
||||
size<3>(grid_shape_up)); // L
|
||||
}
|
||||
|
||||
template<class ProblemShapeMNKL, class TileShape, class AtomThrShape, class ClusterShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_tiled_cta_shape_mnl(ProblemShapeMNKL problem_shape_mnkl,
|
||||
TileShape tile_shape_mnk,
|
||||
AtomThrShape atom_thr_shape_mnk,
|
||||
ClusterShape cluster_shape_mnk) {
|
||||
auto [tiles_m, tiles_n, tiles_l] = product_each(ceil_div(select<0,1,3>(problem_shape_mnkl), take<0,2>(tile_shape_mnk)));
|
||||
auto ctas_m = round_nearest(tiles_m * size<0>(atom_thr_shape_mnk), size<0>(cluster_shape_mnk));
|
||||
auto ctas_n = round_nearest(tiles_n * size<1>(atom_thr_shape_mnk), size<1>(cluster_shape_mnk));
|
||||
auto ctas_l = tiles_l;
|
||||
|
||||
return {static_cast<uint32_t>(ctas_m),
|
||||
static_cast<uint32_t>(ctas_n),
|
||||
static_cast<uint32_t>(ctas_l)};
|
||||
}
|
||||
|
||||
|
||||
// Get clcID and success bit
|
||||
[[nodiscard]] CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_current_work(PipelineState<Stages> state) {
|
||||
uint32_t smem_addr = cute::cast_smem_ptr_to_uint(&clc_response_ptr_[state.index()]);
|
||||
auto work_tile = work_tile_info_from_clc_response(smem_addr);
|
||||
possibly_transpose_work_tile(work_tile);
|
||||
return work_tile;
|
||||
}
|
||||
|
||||
// Set data SMEM ptr
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
set_data_ptr(CLCResponse* clc_response_ptr) {
|
||||
clc_response_ptr_ = clc_response_ptr;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
valid_warpgroup_in_work_tile(WorkTileInfo const& work_tile_info) {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
requires_separate_reduction(Params const& params) {
|
||||
return false;
|
||||
}
|
||||
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE
|
||||
static void
|
||||
fixup(Params const&, WorkTileInfo const&, FrgTensorC&, uint32_t, uint32_t) {}
|
||||
|
||||
|
||||
CUTLASS_DEVICE
|
||||
auto
|
||||
fetch_next_work(WorkTileInfo work_tile_info) {
|
||||
return cute::make_tuple(work_tile_info, true);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static cute::tuple<int32_t, int32_t>
|
||||
possibly_transpose_work_tile(Params::RasterOrder raster_order, int32_t M_idx, int32_t N_idx, FastDivmod divmod_cluster_shape_m, FastDivmod divmod_cluster_shape_n) {
|
||||
if (raster_order == Params::RasterOrder::AlongN) {
|
||||
int cluster_m, remainder_m, cluster_n, remainder_n;
|
||||
divmod_cluster_shape_m(cluster_m, remainder_m, M_idx);
|
||||
divmod_cluster_shape_n(cluster_n, remainder_n, N_idx);
|
||||
M_idx = cluster_n * divmod_cluster_shape_m.divisor + remainder_m;
|
||||
N_idx = cluster_m * divmod_cluster_shape_n.divisor + remainder_n;
|
||||
}
|
||||
return cute::make_tuple(M_idx, N_idx);
|
||||
}
|
||||
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static void
|
||||
possibly_transpose_work_tile(WorkTileInfo& work_tile_info, Params const& params) {
|
||||
auto [M_idx, N_idx] = possibly_transpose_work_tile(
|
||||
params.raster_order_, work_tile_info.M_idx, work_tile_info.N_idx, params.divmod_cluster_shape_m_, params.divmod_cluster_shape_n_);
|
||||
work_tile_info.M_idx = M_idx;
|
||||
work_tile_info.N_idx = N_idx;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
possibly_transpose_work_tile(WorkTileInfo& work_tile_info) {
|
||||
possibly_transpose_work_tile(work_tile_info, scheduler_params);
|
||||
}
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // end namespace cutlass::gemm::kernel::detail
|
||||
+309
@@ -0,0 +1,309 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/barrier.h"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_tile_scheduler_group.hpp"
|
||||
#include "cutlass/gemm/kernel/sm100_tile_scheduler.hpp"
|
||||
#include "cutlass/gemm/kernel/tile_scheduler_params.h"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::kernel::detail {
|
||||
|
||||
//////////////////// Blackwell Grouped Static Scheduler /////////////////////////
|
||||
|
||||
// This tile scheduler is a SM100 wrapper for scheduling by the SM90 Group tile scheduler.
|
||||
// This helps to enable reusing SM90 group tile scheduling capability for SM100 kernels
|
||||
// (e.g., support for CTA rasterization).
|
||||
|
||||
// For Grouped GEMM, most common use case have Problem Shapes for all groups only on device.
|
||||
// Therefore, we don't how many tiles there will be for the scheduler to hand out.
|
||||
// Hence, we have a SM90 style static group scheduler that launches the largest grid possible.
|
||||
// If we had access to host-side problem shapes, one could to use it to figure out the grid shape
|
||||
// and thereafter use CLC query (which can then be linearized and mapped to an approriate tile coord).
|
||||
|
||||
template<class GroupProblemShape>
|
||||
class PersistentTileSchedulerSm100Group {
|
||||
|
||||
public:
|
||||
using UnderlyingScheduler = PersistentTileSchedulerSm90Group<GroupProblemShape>;
|
||||
using UnderlyingProblemShape = typename GroupProblemShape::UnderlyingProblemShape;
|
||||
using Params = PersistentTileSchedulerSm100GroupParams<UnderlyingProblemShape>;
|
||||
using WorkTileInfo = typename UnderlyingScheduler::WorkTileInfo;
|
||||
using Arguments = typename UnderlyingScheduler::Arguments;
|
||||
using RasterOrder = typename Params::RasterOrder;
|
||||
using RasterOrderOptions = typename Params::RasterOrderOptions;
|
||||
struct CLCResponse { uint32_t data[4]; };
|
||||
|
||||
static constexpr bool IsDynamicPersistent = UnderlyingScheduler::IsDynamicPersistent;
|
||||
|
||||
private:
|
||||
UnderlyingScheduler scheduler_sm90;
|
||||
|
||||
public:
|
||||
template <class TileShape, class AtomThrShape, class ClusterShape>
|
||||
static Params
|
||||
to_underlying_arguments(
|
||||
GroupProblemShape problem_shapes,
|
||||
TileShape tile_shape_mnk,
|
||||
AtomThrShape atom_thr_shape_mnk,
|
||||
ClusterShape cluster_shape_mnk,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr) {
|
||||
|
||||
// We only need the tile and cluster shape during scheduler setup, so let FTAD do the magic
|
||||
static_assert(cute::is_static<TileShape>::value);
|
||||
|
||||
auto selected_cluster_shape = cutlass::detail::select_cluster_shape(cluster_shape_mnk, hw_info.cluster_shape);
|
||||
auto cta_shape = cute::conditional_return<not cute::is_static_v<ClusterShape>>(
|
||||
shape_div(tile_shape_mnk, atom_thr_shape_mnk), // Dynamic Cluster: For 2SM kernels, use CTA tile shape for the underlying scheduler
|
||||
shape_div(tile_shape_mnk, selected_cluster_shape)); // Static Cluster: Blackwell builders expects TileShape to be Cluster's Tile Shape, Hopper doesn't
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(
|
||||
problem_shapes.groups(),
|
||||
problem_shapes,
|
||||
hw_info,
|
||||
cta_shape, selected_cluster_shape);
|
||||
|
||||
Params params;
|
||||
params.initialize(
|
||||
problem_blocks,
|
||||
problem_shapes.groups(),
|
||||
problem_shapes.problem_shapes,
|
||||
problem_shapes.host_problem_shapes,
|
||||
to_gemm_coord(cta_shape),
|
||||
to_gemm_coord(selected_cluster_shape),
|
||||
hw_info,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order
|
||||
);
|
||||
|
||||
return params;
|
||||
}
|
||||
|
||||
static bool
|
||||
can_implement(Arguments const& args) {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileSchedulerSm100Group() { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileSchedulerSm100Group(CLCResponse* /* clc_response_ptr */, Params const& params)
|
||||
: scheduler_params(params),
|
||||
scheduler_sm90(params.params_sm90_) { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileSchedulerSm100Group(CLCResponse* /* clc_response_ptr */, Params const& params, dim3 /* block_id_in_cluster */)
|
||||
: scheduler_params(params),
|
||||
scheduler_sm90(params.params_sm90_) { }
|
||||
|
||||
template <class ClusterShape>
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
initial_work_tile_info(ClusterShape cluster_shape) {
|
||||
return scheduler_sm90.initial_work_tile_info(cluster_shape);
|
||||
}
|
||||
|
||||
template<class BlockShape, class ClusterShape>
|
||||
CUTLASS_HOST_DEVICE static
|
||||
dim3
|
||||
get_tiled_cta_shape_mnl(int groups, GroupProblemShape problem_shapes, KernelHardwareInfo hw_info, BlockShape cta_shape, ClusterShape cluster_shape) {
|
||||
return UnderlyingScheduler::get_tiled_cta_shape_mnl(groups, problem_shapes, hw_info, cta_shape, cluster_shape);
|
||||
}
|
||||
|
||||
// Given the inputs, computes the physical grid we should launch.
|
||||
template<class BlockShape, class AtomThrShape, class ClusterShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_grid_shape(
|
||||
Params const& params,
|
||||
GroupProblemShape problem_shapes,
|
||||
BlockShape cta_shape,
|
||||
[[maybe_unused]] AtomThrShape atom_thr_shape,
|
||||
ClusterShape cluster_shape,
|
||||
KernelHardwareInfo hw_info) {
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(
|
||||
problem_shapes.groups(),
|
||||
problem_shapes,
|
||||
hw_info,
|
||||
cta_shape,
|
||||
cluster_shape);
|
||||
|
||||
// Given device SM count, set grid size s.t. we do not launch more thread blocks than we can run concurrently
|
||||
Arguments args{};
|
||||
if constexpr (!std::is_const_v<decltype(args.max_swizzle_size)>) {
|
||||
args.max_swizzle_size = 1 << params.params_sm90_.log_swizzle_size_;
|
||||
}
|
||||
args.raster_order = params.params_sm90_.raster_order_ == RasterOrder::AlongN ? RasterOrderOptions::AlongN : RasterOrderOptions::AlongM;
|
||||
|
||||
return Params::get_grid_shape(
|
||||
problem_blocks,
|
||||
to_gemm_coord(cluster_shape),
|
||||
hw_info,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
/* truncate_by_problem_size = */true,
|
||||
cute::is_static_v<ClusterShape> ? true : false
|
||||
);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static auto
|
||||
work_tile_to_cta_coord(WorkTileInfo work_tile_info) {
|
||||
// SM90 static scheduler implicitly handles CTA coord in a Cluster
|
||||
return make_coord(
|
||||
work_tile_info.M_idx,
|
||||
work_tile_info.N_idx,
|
||||
_,
|
||||
work_tile_info.L_idx
|
||||
);
|
||||
}
|
||||
|
||||
//
|
||||
// K Tile API
|
||||
//
|
||||
template <class ProblemShape, class TileShape, class Shape>
|
||||
CUTLASS_DEVICE
|
||||
auto
|
||||
get_k_tile_iterator(WorkTileInfo const& work_tile_info, ProblemShape problem_shape_MNKL, TileShape tile_shape, Shape) {
|
||||
auto k_tiles = cute::ceil_div(cute::get<2>(problem_shape_MNKL), cute::get<2>(tile_shape));
|
||||
return cute::make_coord_iterator(k_tiles);
|
||||
}
|
||||
|
||||
// Returns whether the block assigned this work should compute the epilogue for the corresponding
|
||||
// output tile. For the Group tile scheduler, this is always true.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
compute_epilogue(WorkTileInfo const&, Params const&) {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
compute_epilogue(WorkTileInfo const&) {
|
||||
return true;
|
||||
}
|
||||
|
||||
// Returns whether fixup is needed for `work_tile_info`. None of the work units returned by
|
||||
// this scheduler require fixup, since none of the work units partition the reduction extent.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
requires_fixup(Params const& params, WorkTileInfo const work_tile_info) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Performs the reduction across splits for a given output tile. No fixup is required for
|
||||
// work units returned by this scheduler.
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
fixup(WorkTileInfo const&, FrgTensorC&, uint32_t, uint32_t, uint32_t = 1) const { }
|
||||
|
||||
template <class ProblemShape, class ElementAccumulator>
|
||||
static size_t
|
||||
get_workspace_size(Arguments const& args, ProblemShape problem_shape, KernelHardwareInfo const& hw_info, uint32_t, uint32_t = 1, uint32_t = 1) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
template <class ElementAccumulator, class ProblemShape, class TileShapeMNK, class AtomThrShape, class ClusterShape>
|
||||
static size_t
|
||||
get_workspace_size(Arguments const& args, ProblemShape problem_shape, TileShapeMNK, AtomThrShape, ClusterShape, KernelHardwareInfo const& hw_info,
|
||||
uint32_t reduction_warp_groups, uint32_t num_accumulator_mtxs = 1) {
|
||||
return 0;
|
||||
}
|
||||
|
||||
template <class ProblemShape, class TileShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int
|
||||
get_work_k_tile_count(WorkTileInfo const& work_tile_info, ProblemShape problem_shape_MNKL, TileShape tile_shape) {
|
||||
// All work units returned by this scheduler cover the entire K iteration
|
||||
// space of the output tile assigned to the work unit.
|
||||
return cute::size(cute::ceil_div(cute::get<2>(problem_shape_MNKL), cute::get<2>(tile_shape)));
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static uint32_t
|
||||
get_work_k_tile_start(WorkTileInfo const&) {
|
||||
// All work units returned by this scheduler start from K tile 0
|
||||
return 0u;
|
||||
}
|
||||
|
||||
template <class ProblemShape, class ElementAccumulator>
|
||||
static cutlass::Status
|
||||
initialize_workspace(Arguments const&, void*, cudaStream_t, ProblemShape const&, KernelHardwareInfo const&, uint32_t, uint32_t = 1, uint32_t = 1, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
|
||||
template <class ElementAccumulator, class ProblemShape, class TileShapeMNK, class AtomThrShape, class ClusterShape>
|
||||
static cutlass::Status
|
||||
initialize_workspace(Arguments const&, void*, cudaStream_t, ProblemShape const&, TileShapeMNK, AtomThrShape, ClusterShape, KernelHardwareInfo const&,
|
||||
uint32_t, uint32_t = 1, CudaHostAdapter *cuda_adapter = nullptr) {
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
|
||||
// Kernel helper function to get next CLC ID
|
||||
template <class CLCPipeline, class CLCPipelineState>
|
||||
CUTLASS_DEVICE
|
||||
auto
|
||||
fetch_next_work(
|
||||
WorkTileInfo work_tile_info,
|
||||
[[maybe_unused]] CLCPipeline& clc_pipeline,
|
||||
[[maybe_unused]] CLCPipelineState clc_pipe_consumer_state) {
|
||||
|
||||
return scheduler_sm90.fetch_next_work(work_tile_info);
|
||||
}
|
||||
|
||||
private:
|
||||
//
|
||||
// Methods
|
||||
//
|
||||
[[nodiscard]] CUTLASS_DEVICE
|
||||
static CLCResponse
|
||||
load_query_response(uint32_t smem_ptr) {
|
||||
return UnderlyingScheduler::load_query_response(smem_ptr);
|
||||
}
|
||||
//
|
||||
// Storage
|
||||
//
|
||||
CLCResponse *clc_response_ptr_ = nullptr;
|
||||
Params scheduler_params;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // end namespace cutlass::gemm::kernel::detail
|
||||
@@ -0,0 +1,979 @@
|
||||
/***************************************************************************************************
|
||||
* 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/arch/barrier.h"
|
||||
#include "cutlass/pipeline/pipeline.hpp"
|
||||
#include "cutlass/gemm/kernel/sm100_tile_scheduler.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_tile_scheduler_stream_k.hpp"
|
||||
#include "cutlass/gemm/kernel/tile_scheduler_params.h"
|
||||
////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::kernel::detail {
|
||||
|
||||
// Persistent Thread Block (TB) scheduler leveraging stream-K decomposition
|
||||
template <
|
||||
class TileShape,
|
||||
class ClusterShape,
|
||||
uint32_t Stages_
|
||||
>
|
||||
class PersistentTileSchedulerSm100StreamK {
|
||||
using UnderlyingScheduler = PersistentTileSchedulerSm100<ClusterShape, Stages_>;
|
||||
using UnderlyingStreamKScheduler = PersistentTileSchedulerSm90StreamK<TileShape, ClusterShape>;
|
||||
using InternalWorkTileInfo = typename UnderlyingScheduler::WorkTileInfo;
|
||||
using InternalParams = typename UnderlyingScheduler::Params;
|
||||
// Shapediv failures currently occur with tile shape N of 192
|
||||
static constexpr bool ForceDataParallel = size<1>(TileShape{}) == 192;
|
||||
|
||||
public:
|
||||
static constexpr uint32_t Stages = Stages_;
|
||||
|
||||
using CLCResponse = typename UnderlyingScheduler::CLCResponse;
|
||||
using WorkTileInfo = typename UnderlyingStreamKScheduler::WorkTileInfo;
|
||||
using Arguments = typename UnderlyingStreamKScheduler::Arguments;
|
||||
|
||||
using Params = PersistentTileSchedulerSm100StreamKParams;
|
||||
using RasterOrder = PersistentTileSchedulerSm90Params::RasterOrder;
|
||||
using RasterOrderOptions = PersistentTileSchedulerSm90Params::RasterOrderOptions;
|
||||
|
||||
using SharedStorage = typename UnderlyingScheduler::SharedStorage;
|
||||
using Pipeline = typename UnderlyingScheduler::Pipeline;
|
||||
using ThrottlePipeline = typename UnderlyingScheduler::ThrottlePipeline;
|
||||
|
||||
static constexpr bool IsDynamicPersistent = true;
|
||||
|
||||
// Number of sub blocks in the kernel epilogue
|
||||
static constexpr int EpilogueSubtiles = 1;
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
PersistentTileSchedulerSm100StreamK() { }
|
||||
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileSchedulerSm100StreamK(Params const& params)
|
||||
: sm100_scheduler_(params.sm100_params_)
|
||||
, params_(params)
|
||||
, block_id_in_cluster_(cute::block_id_in_cluster()) {
|
||||
// Set the current linear idx to be equal to the linear idx of the first work tile to be computed
|
||||
auto cs = make_shape(
|
||||
params.sm100_params_.divmod_cluster_shape_m_.divisor,
|
||||
params.sm100_params_.divmod_cluster_shape_n_.divisor,
|
||||
Int<1>{});
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileSchedulerSm100StreamK(CLCResponse* clc_response_ptr, Params const& params, dim3 block_id_in_cluster)
|
||||
: sm100_scheduler_(clc_response_ptr, params.sm100_params_, block_id_in_cluster),
|
||||
params_(params),
|
||||
block_id_in_cluster_(block_id_in_cluster) {
|
||||
// Set the current linear idx to be equal to the linear idx of the first work tile to be computed
|
||||
auto cs = make_shape(
|
||||
params.sm100_params_.divmod_cluster_shape_m_.divisor,
|
||||
params.sm100_params_.divmod_cluster_shape_n_.divisor,
|
||||
Int<1>{});
|
||||
}
|
||||
|
||||
template <class ProblemShape, class TileShapeMNK>
|
||||
CUTLASS_DEVICE
|
||||
PersistentTileSchedulerSm100StreamK(CLCResponse* clc_response_ptr, Params const& params,
|
||||
ProblemShape problem_shape_mnkl, TileShapeMNK tile_shape, dim3 block_id_in_cluster)
|
||||
: PersistentTileSchedulerSm100StreamK(clc_response_ptr, params, block_id_in_cluster) { }
|
||||
|
||||
template <class ProblemShape>
|
||||
static Params
|
||||
to_underlying_arguments(
|
||||
ProblemShape problem_shape,
|
||||
TileShape tile_shape,
|
||||
[[maybe_unused]] ClusterShape cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
Arguments const& args,
|
||||
void* workspace,
|
||||
[[maybe_unused]] const uint32_t epilogue_subtile = 1,
|
||||
uint32_t ktile_start_alignment_count = 1u) {
|
||||
|
||||
auto cs = cutlass::detail::select_cluster_shape(cluster_shape, hw_info.cluster_shape);
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape, 1);
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape, cs);
|
||||
uint32_t k_tile_per_output_tile = cute::size(cute::ceil_div(cute::shape<2>(problem_shape_mnkl), cute::shape<2>(TileShape{})));
|
||||
|
||||
Params params;
|
||||
params.initialize(
|
||||
problem_blocks,
|
||||
k_tile_per_output_tile,
|
||||
to_gemm_coord(cs),
|
||||
hw_info,
|
||||
args.splits,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
args.reduction_mode,
|
||||
ForceDataParallel ? Params::DecompositionMode::DataParallel : args.decomposition_mode,
|
||||
workspace,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
return params;
|
||||
}
|
||||
|
||||
template <class ProblemShape, class TileShapeMNK, class AtomThrShape>
|
||||
static Params
|
||||
to_underlying_arguments(
|
||||
ProblemShape problem_shape_mnkl,
|
||||
TileShapeMNK tile_shape_mnk,
|
||||
AtomThrShape atom_thr_shape_mnk,
|
||||
ClusterShape cluster_shape_mnk,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
Arguments const& args,
|
||||
void* workspace = nullptr,
|
||||
uint32_t ktile_start_alignment_count = 1u
|
||||
) {
|
||||
|
||||
auto cs = cutlass::detail::select_cluster_shape(cluster_shape_mnk, hw_info.cluster_shape);
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape_mnk, atom_thr_shape_mnk, cs);
|
||||
uint32_t k_tile_per_output_tile = cute::size(cute::ceil_div(cute::shape<2>(problem_shape_mnkl), cute::shape<2>(TileShape{})));
|
||||
|
||||
Params params;
|
||||
params.initialize(
|
||||
problem_blocks,
|
||||
k_tile_per_output_tile,
|
||||
to_gemm_coord(cs),
|
||||
hw_info,
|
||||
args.splits,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
args.reduction_mode,
|
||||
ForceDataParallel ? Params::DecompositionMode::DataParallel : args.decomposition_mode,
|
||||
workspace,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
|
||||
return params;
|
||||
}
|
||||
|
||||
static bool
|
||||
can_implement(Arguments const& args) {
|
||||
return UnderlyingStreamKScheduler::can_implement(args);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
PipelineState<Stages>
|
||||
advance_to_next_work(Pipeline& clc_pipeline, PipelineState<Stages> clc_pipe_producer_state) const {
|
||||
return sm100_scheduler_.advance_to_next_work(clc_pipeline, clc_pipe_producer_state);
|
||||
}
|
||||
|
||||
// Get clcID and success bit
|
||||
[[nodiscard]] CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
get_current_work(PipelineState<Stages> state) {
|
||||
InternalWorkTileInfo work_tile_info = sm100_scheduler_.get_current_work(state);
|
||||
if (!work_tile_info.is_valid()) {
|
||||
return invalid_work_tile();
|
||||
}
|
||||
|
||||
return convert_work(work_tile_info);
|
||||
}
|
||||
// Given the inputs, computes the total number of output blocks this problem will compute over
|
||||
template<class ProblemShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_tiled_cta_shape_mnl(ProblemShape problem_shape_mnkl, TileShape blk_shape, ClusterShape cluster_shape) {
|
||||
return UnderlyingScheduler::get_tiled_cta_shape_mnl(problem_shape_mnkl, blk_shape, cluster_shape);
|
||||
}
|
||||
|
||||
template<class ProblemShape, class TileShapeMNK, class AtomThrShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_tiled_cta_shape_mnl(ProblemShape problem_shape_mnkl,
|
||||
TileShapeMNK tile_shape_mnk,
|
||||
AtomThrShape atom_thr_shape_mnk,
|
||||
ClusterShape cluster_shape_mnk) {
|
||||
return UnderlyingScheduler::get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape_mnk, atom_thr_shape_mnk, cluster_shape_mnk);
|
||||
}
|
||||
|
||||
// Given the inputs, computes the physical grid we should launch.
|
||||
template <class ProblemShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_grid_shape(
|
||||
Params const& params,
|
||||
ProblemShape problem_shape,
|
||||
TileShape tile_shape,
|
||||
ClusterShape cluster_shape,
|
||||
KernelHardwareInfo hw_info,
|
||||
[[maybe_unused]] Arguments arguments) {
|
||||
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape, 1);
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape, cluster_shape);
|
||||
return params.get_grid_shape(problem_blocks, to_gemm_coord(cluster_shape));
|
||||
}
|
||||
|
||||
// Given the inputs, computes the physical grid we should launch.
|
||||
template<class ProblemShape, class TileShapeMNK, class AtomThrShape>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_grid_shape(
|
||||
Params const& params,
|
||||
ProblemShape problem_shape_mnkl,
|
||||
TileShapeMNK tile_shape_mnk,
|
||||
AtomThrShape atom_thr_shape_mnk,
|
||||
ClusterShape cluster_shape_mnk,
|
||||
KernelHardwareInfo hw_info) {
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape_mnk, atom_thr_shape_mnk, cluster_shape_mnk);
|
||||
return params.get_grid_shape(problem_blocks, to_gemm_coord(cluster_shape_mnk));
|
||||
}
|
||||
|
||||
|
||||
// Returns the initial work tile info that will be computed over
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
initial_work_tile_info(ClusterShape cluster_shape) {
|
||||
InternalWorkTileInfo work_tile_info = UnderlyingScheduler::initial_work_tile_info(cluster_shape, params_.sm100_params_);
|
||||
work_tile_info.is_valid_tile = false;
|
||||
return convert_work(work_tile_info);
|
||||
}
|
||||
|
||||
// Returns a CTA-tiled coordinate for the provided work tile info
|
||||
CUTLASS_DEVICE
|
||||
auto
|
||||
work_tile_to_cta_coord(WorkTileInfo const& work_tile_info) {
|
||||
if (is_dp_only()) {
|
||||
// For data-parallel decompositions, simply default to the
|
||||
// underlying SM100 scheduler.
|
||||
auto underlying_work_tile = to_underlying_work_tile_info(work_tile_info);
|
||||
return sm100_scheduler_.work_tile_to_cta_coord(underlying_work_tile);
|
||||
}
|
||||
else {
|
||||
// The SM90 stream-K scheduler already operates only at CTA level,
|
||||
// so the returned work tile info already contains CTA offsets within
|
||||
// each cluster tile.
|
||||
return cute::make_coord(
|
||||
work_tile_info.M_idx,
|
||||
work_tile_info.N_idx,
|
||||
_,
|
||||
work_tile_info.L_idx
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
// Returns whether the current work_tile_info passed in should continue to be used.
|
||||
CUTLASS_DEVICE
|
||||
bool
|
||||
continue_current_work(WorkTileInfo& work_tile_info) const {
|
||||
return UnderlyingStreamKScheduler::continue_current_work_for_linear_idx(
|
||||
current_work_linear_idx_, unit_iter_start_, block_id_in_cluster_, work_tile_info, params_.sk_params_);
|
||||
}
|
||||
|
||||
// Kernel helper function to get next CLC ID and whether to advance the CLC pipeline state.
|
||||
template <class CLCPipeline, class CLCPipelineState>
|
||||
CUTLASS_DEVICE
|
||||
cute::tuple<WorkTileInfo, bool>
|
||||
fetch_next_work(
|
||||
WorkTileInfo work_tile_info,
|
||||
CLCPipeline& clc_pipeline,
|
||||
CLCPipelineState clc_pipe_consumer_state) {
|
||||
// Check whether we should continue on with the current work unit. If this is the case,
|
||||
// the work unit will have been updated in continue_current_work to reflect the new
|
||||
// tile to be computed. Return `false` to indicate that the CLC pipeline state
|
||||
// need not be advanced.
|
||||
if (continue_current_work(work_tile_info)) {
|
||||
return cute::make_tuple(work_tile_info, false);
|
||||
}
|
||||
|
||||
clc_pipeline.consumer_wait(clc_pipe_consumer_state);
|
||||
auto new_work_tile_info = get_current_work(clc_pipe_consumer_state);
|
||||
clc_pipeline.consumer_release(clc_pipe_consumer_state);
|
||||
|
||||
// Return true to indicate that the CLC pipeline state should be advanced
|
||||
return cute::make_tuple(new_work_tile_info, true);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
cute::tuple<WorkTileInfo, bool>
|
||||
fetch_next_work(WorkTileInfo work_tile_info) {
|
||||
return cute::make_tuple(work_tile_info, true);
|
||||
}
|
||||
|
||||
// Set data SMEM ptr
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
set_data_ptr(CLCResponse* clc_response_ptr) {
|
||||
sm100_scheduler_.set_data_ptr(clc_response_ptr);
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
valid_warpgroup_in_work_tile(WorkTileInfo const& work_tile_info) {
|
||||
return true;
|
||||
}
|
||||
|
||||
CUTLASS_DEVICE
|
||||
static bool
|
||||
requires_separate_reduction(Params const& params) {
|
||||
return false;
|
||||
}
|
||||
|
||||
// Returns whether the block assigned this work should compute the epilogue for the corresponding
|
||||
// output tile. For the case of stream-K, this should only occur if the work is marked as the final split.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static bool
|
||||
compute_epilogue(WorkTileInfo const& work_tile_info, Params const& params) {
|
||||
return UnderlyingStreamKScheduler::compute_epilogue(work_tile_info, params.sk_params_);
|
||||
}
|
||||
|
||||
// Non-static variant of compute_epilogue. Used in cases where passing
|
||||
// in Params is inconvenient.
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
compute_epilogue(WorkTileInfo const& work_tile_info) const {
|
||||
return UnderlyingStreamKScheduler::compute_epilogue(work_tile_info, params_.sk_params_);
|
||||
}
|
||||
|
||||
template <class ProblemShape, class ElementAccumulator>
|
||||
static size_t
|
||||
get_workspace_size(
|
||||
Arguments const& args,
|
||||
ProblemShape problem_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
uint32_t reduction_warp_groups,
|
||||
[[maybe_unused]] const uint32_t epilogue_subtile = 1,
|
||||
uint32_t num_accumulator_mtxs = 1,
|
||||
uint32_t ktile_start_alignment_count = 1) {
|
||||
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape, 1);
|
||||
|
||||
auto cs = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape);
|
||||
TileShape tile_shape;
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape, cs);
|
||||
uint32_t k_tile_per_output_tile = cute::size(cute::ceil_div(cute::shape<2>(problem_shape_mnkl), cute::shape<2>(TileShape{})));
|
||||
|
||||
return Params::get_workspace_size(
|
||||
problem_blocks,
|
||||
k_tile_per_output_tile,
|
||||
to_gemm_coord(tile_shape),
|
||||
to_gemm_coord(cs),
|
||||
hw_info,
|
||||
args.splits,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
ForceDataParallel ? Params::DecompositionMode::DataParallel : args.decomposition_mode,
|
||||
args.reduction_mode,
|
||||
reduction_warp_groups,
|
||||
sizeof_bits<typename UnderlyingStreamKScheduler::BarrierType>::value,
|
||||
sizeof_bits<ElementAccumulator>::value,
|
||||
EpilogueSubtiles,
|
||||
num_accumulator_mtxs,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
template <class ElementAccumulator, class ProblemShape, class TileShapeMNK, class AtomThrShape>
|
||||
static size_t
|
||||
get_workspace_size(
|
||||
Arguments const& args,
|
||||
ProblemShape problem_shape,
|
||||
TileShapeMNK tile_shape_mnk,
|
||||
AtomThrShape atom_thr_shape_mnk,
|
||||
ClusterShape cluster_shape_mnk,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
uint32_t reduction_warp_groups,
|
||||
uint32_t num_accumulator_mtxs = 1,
|
||||
uint32_t ktile_start_alignment_count = 1) {
|
||||
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape, 1);
|
||||
|
||||
auto cs = cutlass::detail::select_cluster_shape(cluster_shape_mnk, hw_info.cluster_shape);
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape_mnk, atom_thr_shape_mnk, cs);
|
||||
uint32_t k_tile_per_output_tile = cute::size(cute::ceil_div(cute::shape<2>(problem_shape_mnkl), cute::shape<2>(TileShape{})));
|
||||
|
||||
auto cta_tile_shape_mnk = shape_div(tile_shape_mnk, atom_thr_shape_mnk);
|
||||
|
||||
return Params::get_workspace_size(
|
||||
problem_blocks,
|
||||
k_tile_per_output_tile,
|
||||
to_gemm_coord(cta_tile_shape_mnk),
|
||||
to_gemm_coord(cs),
|
||||
hw_info,
|
||||
args.splits,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
ForceDataParallel ? Params::DecompositionMode::DataParallel : args.decomposition_mode,
|
||||
args.reduction_mode,
|
||||
reduction_warp_groups,
|
||||
sizeof_bits<typename UnderlyingStreamKScheduler::BarrierType>::value,
|
||||
sizeof_bits<ElementAccumulator>::value,
|
||||
EpilogueSubtiles,
|
||||
num_accumulator_mtxs,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
template <class ProblemShape, class ElementAccumulator>
|
||||
static cutlass::Status
|
||||
initialize_workspace(
|
||||
Arguments const& args,
|
||||
void* workspace,
|
||||
cudaStream_t stream,
|
||||
ProblemShape const& problem_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
uint32_t reduction_warp_groups,
|
||||
[[maybe_unused]] const uint32_t epilogue_subtile = 1,
|
||||
uint32_t num_accumulator_mtxs = 1,
|
||||
CudaHostAdapter *cuda_adapter = nullptr,
|
||||
uint32_t ktile_start_alignment_count = 1) {
|
||||
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape, 1);
|
||||
|
||||
auto cs = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape);
|
||||
TileShape tile_shape;
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape, cs);
|
||||
uint32_t k_tile_per_output_tile = cute::size(cute::ceil_div(cute::shape<2>(problem_shape_mnkl), cute::shape<2>(TileShape{})));
|
||||
|
||||
return Params::initialize_workspace(
|
||||
workspace,
|
||||
stream,
|
||||
problem_blocks,
|
||||
k_tile_per_output_tile,
|
||||
to_gemm_coord(tile_shape),
|
||||
to_gemm_coord(cs),
|
||||
hw_info,
|
||||
args.splits,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
ForceDataParallel ? Params::DecompositionMode::DataParallel : args.decomposition_mode,
|
||||
args.reduction_mode,
|
||||
reduction_warp_groups,
|
||||
sizeof_bits<typename UnderlyingStreamKScheduler::BarrierType>::value,
|
||||
sizeof_bits<ElementAccumulator>::value,
|
||||
EpilogueSubtiles,
|
||||
num_accumulator_mtxs,
|
||||
cuda_adapter,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
template <class ElementAccumulator, class ProblemShape, class TileShapeMNK, class AtomThrShape>
|
||||
static cutlass::Status
|
||||
initialize_workspace(
|
||||
Arguments const& args,
|
||||
void* workspace,
|
||||
cudaStream_t stream,
|
||||
ProblemShape const& problem_shape,
|
||||
TileShapeMNK tile_shape_mnk,
|
||||
AtomThrShape atom_thr_shape_mnk,
|
||||
ClusterShape cluster_shape_mnk,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
uint32_t reduction_warp_groups,
|
||||
uint32_t num_accumulator_mtxs = 1,
|
||||
CudaHostAdapter *cuda_adapter = nullptr,
|
||||
uint32_t ktile_start_alignment_count = 1) {
|
||||
|
||||
auto problem_shape_mnkl = cute::append<4>(problem_shape, 1);
|
||||
|
||||
auto cs = cutlass::detail::select_cluster_shape(cluster_shape_mnk, hw_info.cluster_shape);
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape_mnkl, tile_shape_mnk, atom_thr_shape_mnk, cs);
|
||||
uint32_t k_tile_per_output_tile = cute::size(cute::ceil_div(cute::shape<2>(problem_shape_mnkl), cute::shape<2>(TileShape{})));
|
||||
|
||||
auto cta_tile_shape_mnk = shape_div(tile_shape_mnk, atom_thr_shape_mnk);
|
||||
|
||||
return Params::initialize_workspace(
|
||||
workspace,
|
||||
stream,
|
||||
problem_blocks,
|
||||
k_tile_per_output_tile,
|
||||
to_gemm_coord(cta_tile_shape_mnk),
|
||||
to_gemm_coord(cs),
|
||||
hw_info,
|
||||
args.splits,
|
||||
args.max_swizzle_size,
|
||||
args.raster_order,
|
||||
ForceDataParallel ? Params::DecompositionMode::DataParallel : args.decomposition_mode,
|
||||
args.reduction_mode,
|
||||
reduction_warp_groups,
|
||||
sizeof_bits<typename UnderlyingStreamKScheduler::BarrierType>::value,
|
||||
sizeof_bits<ElementAccumulator>::value,
|
||||
EpilogueSubtiles,
|
||||
num_accumulator_mtxs,
|
||||
cuda_adapter,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
template <class ProblemShape, class TileShapeMNK>
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int
|
||||
get_work_k_tile_count(WorkTileInfo const& work_tile_info, ProblemShape, TileShapeMNK) {
|
||||
return work_tile_info.k_tile_count;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static uint32_t
|
||||
get_work_k_tile_start(WorkTileInfo const& work_tile_info) {
|
||||
return work_tile_info.K_idx;
|
||||
}
|
||||
|
||||
template <class ProblemShape, class TileShapeMNK, class Shape>
|
||||
CUTLASS_DEVICE
|
||||
auto
|
||||
get_k_tile_iterator(WorkTileInfo const& work_tile_info, ProblemShape problem_shape, TileShapeMNK tile_shape, Shape) {
|
||||
// Get the shape of k tiles instead of the counter. Otherwise, if the problem shape has
|
||||
// multiple k modes, the DMA loop would need to decompose the iterator onto every mode
|
||||
// every time global loading happens. This would incur extra overhead.
|
||||
auto k_tiles = cute::ceil_div(cute::get<2>(problem_shape), cute::get<2>(tile_shape));
|
||||
auto k_tile_start = get_work_k_tile_start(work_tile_info);
|
||||
// Iterate start from current k tile start over the k tiles shape.
|
||||
return cute::make_coord_iterator(idx2crd(k_tile_start, k_tiles), k_tiles);
|
||||
}
|
||||
|
||||
// Returns whether fixup is needed for `work_tile_info`.
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
requires_fixup(WorkTileInfo const work_tile_info) const {
|
||||
return UnderlyingStreamKScheduler::requires_fixup(params_.sk_params_, work_tile_info);
|
||||
}
|
||||
|
||||
// Performs the reduction across splits for a given output tile.
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE
|
||||
void
|
||||
fixup(
|
||||
WorkTileInfo const& work_tile_info,
|
||||
FrgTensorC& accumulators,
|
||||
uint32_t num_barriers,
|
||||
uint32_t barrier_idx,
|
||||
uint32_t num_accumulator_mtxs = 1) const {
|
||||
|
||||
using BarrierManager = SyncManager<cutlass::detail::SyncwarpSync, NumThreadsPerWarp>;
|
||||
|
||||
UnderlyingStreamKScheduler s;
|
||||
return s.template fixup_helper<FrgTensorC, BarrierManager>(
|
||||
params_.sk_params_, work_tile_info, accumulators, num_barriers, barrier_idx, num_accumulator_mtxs);
|
||||
}
|
||||
|
||||
|
||||
// Performs the reduction across splits for a given output tile.
|
||||
template <class FrgTensorC>
|
||||
CUTLASS_DEVICE
|
||||
static void
|
||||
fixup(
|
||||
Params const& params,
|
||||
WorkTileInfo const& work_tile_info,
|
||||
FrgTensorC& accumulators,
|
||||
uint32_t num_barriers,
|
||||
uint32_t barrier_idx) {
|
||||
UnderlyingStreamKScheduler::fixup(params.sk_params_, work_tile_info, accumulators, num_barriers, barrier_idx);
|
||||
}
|
||||
|
||||
// Performs reduction across splits for a given output tile
|
||||
template <
|
||||
bool IsComplex,
|
||||
class TiledMma,
|
||||
class AccEngine,
|
||||
class AccLayout,
|
||||
class AccumulatorPipeline,
|
||||
class AccumulatorPipelineState,
|
||||
class CopyOpT2R
|
||||
>
|
||||
CUTLASS_DEVICE
|
||||
AccumulatorPipelineState
|
||||
fixup(
|
||||
TiledMma const& tiled_mma,
|
||||
WorkTileInfo const& work_tile_info,
|
||||
cute::Tensor<AccEngine, AccLayout>& accumulators,
|
||||
AccumulatorPipeline acc_pipeline,
|
||||
AccumulatorPipelineState acc_pipe_consumer_state,
|
||||
CopyOpT2R) const {
|
||||
using namespace cute;
|
||||
static_assert(cute::is_rmem_v<AccEngine> || cute::is_tmem_v<AccEngine>, "Accumulator must be in either TMEM or RF");
|
||||
|
||||
if constexpr (ForceDataParallel) {
|
||||
return acc_pipe_consumer_state;
|
||||
}
|
||||
else {
|
||||
if (!requires_fixup(work_tile_info)) {
|
||||
if constexpr (cute::is_tmem_v<AccEngine>) {
|
||||
if (!work_tile_info.is_valid()) {
|
||||
// The first work tile can be invalid, but still must release TMEM
|
||||
acc_pipeline.consumer_wait(acc_pipe_consumer_state);
|
||||
acc_pipeline.consumer_release(acc_pipe_consumer_state);
|
||||
++acc_pipe_consumer_state;
|
||||
}
|
||||
}
|
||||
return acc_pipe_consumer_state;
|
||||
}
|
||||
|
||||
if constexpr (cute::is_tmem_v<AccEngine>) {
|
||||
// When accumulators reside in TMEM, perform TMEM -> RF loads before performing fixup,
|
||||
// and perform RF -> TMEM stores after fixup (when the split must compute the epilogue)
|
||||
if constexpr (IsComplex) {
|
||||
constexpr uint32_t NumAccumulatorMtx = 2;
|
||||
Tensor accumulators_real = accumulators(_,_,_,0);
|
||||
tmem_fixup(
|
||||
tiled_mma,
|
||||
work_tile_info,
|
||||
accumulators_real,
|
||||
acc_pipeline,
|
||||
acc_pipe_consumer_state,
|
||||
CopyOpT2R{},
|
||||
NumAccumulatorMtx,
|
||||
0 /*idx_accumulator_mtx*/
|
||||
);
|
||||
|
||||
Tensor accumulators_imag = accumulators(_,_,_,1);
|
||||
return tmem_fixup(
|
||||
tiled_mma,
|
||||
work_tile_info,
|
||||
accumulators_imag,
|
||||
acc_pipeline,
|
||||
acc_pipe_consumer_state,
|
||||
CopyOpT2R{},
|
||||
NumAccumulatorMtx,
|
||||
1 /*idx_accumulator_mtx*/
|
||||
);
|
||||
}
|
||||
else {
|
||||
return tmem_fixup(
|
||||
tiled_mma,
|
||||
work_tile_info,
|
||||
accumulators,
|
||||
acc_pipeline,
|
||||
acc_pipe_consumer_state,
|
||||
CopyOpT2R{}
|
||||
);
|
||||
}
|
||||
}
|
||||
else {
|
||||
// Simply perform fixup without TMEM loads when accumulators reside in RF
|
||||
constexpr uint32_t ThreadsForFixup = NumThreadsPerWarpGroup;
|
||||
constexpr uint32_t Offset = static_cast<int>(cutlass::arch::ReservedNamedBarriers::StreamkBarrier0);
|
||||
constexpr uint32_t MaxNumNamedBarriers = 1;
|
||||
constexpr uint32_t BarrierIdx = 0;
|
||||
using BarrierManager = NamedBarrierManager<ThreadsForFixup, Offset, MaxNumNamedBarriers>;
|
||||
constexpr int NumAccumulatorMtx = IsComplex ? 2 : 1;
|
||||
|
||||
UnderlyingStreamKScheduler::template fixup_helper<cute::remove_cvref_t<decltype(accumulators)>, BarrierManager>(
|
||||
params_.sk_params_, work_tile_info, accumulators, MaxNumNamedBarriers, BarrierIdx, NumAccumulatorMtx);
|
||||
return acc_pipe_consumer_state;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
|
||||
// Convert CTA-level work tile info to cluster-level tile coord
|
||||
CUTLASS_DEVICE
|
||||
auto
|
||||
work_tile_to_cluster_coord_mnkl(WorkTileInfo work_tile_info) const {
|
||||
typename UnderlyingScheduler::WorkTileInfo tmp{
|
||||
work_tile_info.M_idx,
|
||||
work_tile_info.N_idx,
|
||||
work_tile_info.L_idx,
|
||||
work_tile_info.is_valid()
|
||||
};
|
||||
return sm100_scheduler_.work_tile_to_cluster_coord_mnkl(tmp);
|
||||
}
|
||||
|
||||
private:
|
||||
CUTLASS_HOST_DEVICE
|
||||
WorkTileInfo invalid_work_tile() const {
|
||||
// Mark the work tile as invalid based on its having a 0 K tiles to comptue.
|
||||
// Set the M, N, and L indices to be outside of the range of valid tiles for the problem.
|
||||
return {
|
||||
static_cast<int32_t>(params_.sm100_params_.problem_tiles_m_) * params_.sm100_params_.divmod_cluster_shape_m_.divisor,
|
||||
static_cast<int32_t>(params_.sm100_params_.problem_tiles_n_) * params_.sm100_params_.divmod_cluster_shape_n_.divisor,
|
||||
0, // K_idx
|
||||
static_cast<int32_t>(params_.sm100_params_.problem_tiles_l_),
|
||||
0 // k_tile_count
|
||||
};
|
||||
}
|
||||
|
||||
// Converts the work tile info returned by the SM100 scheduler to a linear index
|
||||
CUTLASS_DEVICE
|
||||
uint64_t
|
||||
to_linear_idx(
|
||||
InternalWorkTileInfo const& work_tile_info,
|
||||
Params const& params) {
|
||||
// The InternalWorkTileInfo returned from CLC query gives all CTAs in a cluster
|
||||
// the tile offset corresponding to the first CTA tile in the cluster tile assigned
|
||||
// to the cluster. Since the SM90 tile scheduler operates at CTA level, we must assign
|
||||
// each CTA its own tile when computing the linear ID to be used by the SM90
|
||||
// stream-K scheduler.
|
||||
auto start_cta_m_preferred_cluster = params.sk_params_.truncate_to_cluster_size_m(work_tile_info.M_idx);
|
||||
auto start_cta_n_preferred_cluster = params.sk_params_.truncate_to_cluster_size_n(work_tile_info.N_idx);
|
||||
uint64_t cluster_idx = gridDim.y * start_cta_m_preferred_cluster + start_cta_n_preferred_cluster;
|
||||
uint64_t sm_count = gridDim.x * gridDim.y;
|
||||
uint64_t wave_idx = work_tile_info.L_idx;
|
||||
|
||||
auto cluster_start_linear_id = sm_count * wave_idx + cluster_idx;
|
||||
|
||||
// Determine the offset of this CTA in the preferred cluster shape.
|
||||
// This calculation aims to accomodate both cases in which this CTA is part of a preferred cluster
|
||||
// and those in which it is part of a fallback cluster.
|
||||
//
|
||||
// The calculation is performed by computing the starting M and N index of the preferred cluster that
|
||||
// this CTA would be in, and then subtracting these from the true CTA M and N indexes.
|
||||
//
|
||||
// In the case where this CTA is part of a preferred cluster, the resulting offsets are equivalent
|
||||
// to those returned by cute::block_id_in_cluster();
|
||||
auto [cta_m_in_cluster, cta_n_in_cluster, _] = block_id_in_cluster_;
|
||||
uint64_t cta_m_in_preferred_cluster = work_tile_info.M_idx + cta_m_in_cluster - start_cta_m_preferred_cluster;
|
||||
uint64_t cta_n_in_preferred_cluster = work_tile_info.N_idx + cta_n_in_cluster - start_cta_n_preferred_cluster;
|
||||
|
||||
if (params.sk_params_.raster_order_ == RasterOrder::AlongN) {
|
||||
return cluster_start_linear_id + (params.sk_params_.divmod_cluster_shape_minor_.divisor * cta_n_in_preferred_cluster) + cta_m_in_preferred_cluster;
|
||||
}
|
||||
else {
|
||||
return cluster_start_linear_id + (params.sk_params_.divmod_cluster_shape_minor_.divisor * cta_m_in_preferred_cluster) + cta_n_in_preferred_cluster;
|
||||
}
|
||||
}
|
||||
|
||||
// Converts the work tile info returned by the SM100 scheduler to a stream-K work tile info
|
||||
CUTLASS_DEVICE
|
||||
WorkTileInfo
|
||||
convert_work(InternalWorkTileInfo const& work_tile_info) {
|
||||
if (has_sk_work()) {
|
||||
current_work_linear_idx_ = to_linear_idx(work_tile_info, params_);
|
||||
auto work = UnderlyingStreamKScheduler::get_current_work_for_linear_idx(unit_iter_start_, current_work_linear_idx_, block_id_in_cluster_, params_.sk_params_);
|
||||
if (!work.is_valid()) {
|
||||
return invalid_work_tile();
|
||||
}
|
||||
return work;
|
||||
}
|
||||
else if (is_split_k()) {
|
||||
// Split-K offsets are returned directly by CLC query (rather than being
|
||||
// returned by the SM90 stream-K tile scheduler). CLC query returns
|
||||
// the first CTA tile of work for each CTA in a cluster, but later use of the
|
||||
// split-K work tile for fixup expect a CTA-offset tile. Thus, we need to offset
|
||||
// each CTA's M and N index by the CTA offset in the cluster.
|
||||
auto [cta_m_in_cluster, cta_n_in_cluster, _] = block_id_in_cluster_;
|
||||
auto M_idx = work_tile_info.M_idx + cta_m_in_cluster;
|
||||
auto N_idx = work_tile_info.N_idx + cta_n_in_cluster;
|
||||
|
||||
int L_idx, Split_idx;
|
||||
params_.sk_params_.divmod_splits_(L_idx, Split_idx, work_tile_info.L_idx);
|
||||
|
||||
// TODO: Modularize the SM90 scheduler to pull out and reuse this redundant code
|
||||
int additional_k_tiles = 0;
|
||||
int split_start_offset = params_.sk_params_.big_units_;
|
||||
|
||||
if (Split_idx < params_.sk_params_.big_units_) {
|
||||
// Offsets for "big" units. One additional k iteration is performed,
|
||||
// and each split preceding us was a big unit, so we must increase
|
||||
// our split starting offset by our split ID (Split_idx).
|
||||
additional_k_tiles = 1;
|
||||
split_start_offset = Split_idx;
|
||||
}
|
||||
|
||||
// Set up k iteration count and split starting iteration assuming the
|
||||
// iteration space is evenly split.
|
||||
uint32_t k_tiles = params_.sk_params_.divmod_k_tiles_per_sk_unit_.divisor;
|
||||
uint32_t K_idx = Split_idx * k_tiles;
|
||||
|
||||
// Apply any fixup needed to handle residuals
|
||||
K_idx += split_start_offset;
|
||||
k_tiles += additional_k_tiles;
|
||||
|
||||
// K_idx is even for each cta.
|
||||
//
|
||||
// * Example
|
||||
// 53 k_tiles per output tile
|
||||
// 10 k_tiles for normal size split
|
||||
// 11 k_tiles for start three big unit
|
||||
//
|
||||
// split 0 : K_idx = [0, 10], k_tiles = 11 -> K_idx = [0, 11], k_tiles = 12
|
||||
// split 1 : K_idx = [11, 21], k_tiles = 11 -> K_idx = [12, 21], k_tiles = 10
|
||||
// split 2 : K_idx = [22, 32], k_tiles = 11 -> K_idx = [22, 33], k_tiles = 12
|
||||
// split 3 : K_idx = [33, 42], k_tiles = 10 -> K_idx = [34, 42], k_tiles = 9 -> K_idx = [34, 43], k_tiles = 10
|
||||
// split 4 : K_idx = [43, 52], k_tiles = 10 -> K_idx = [44, 52], k_tiles = 9
|
||||
if (params_.sk_params_.ktile_start_alignment_count_ == 2u && K_idx % 2 != 0) {
|
||||
// If current cta K_idx not start from even, give up one k_tile
|
||||
K_idx += 1;
|
||||
k_tiles -= 1;
|
||||
}
|
||||
if (params_.sk_params_.ktile_start_alignment_count_ == 2u &&
|
||||
(K_idx + k_tiles) % 2 != 0 &&
|
||||
(K_idx + k_tiles) < params_.sk_params_.divmod_tiles_per_output_tile_.divisor) {
|
||||
// If next cta K_idx not start from even, acquire one k_tile
|
||||
k_tiles += 1;
|
||||
}
|
||||
|
||||
return {
|
||||
static_cast<int32_t>(M_idx),
|
||||
static_cast<int32_t>(N_idx),
|
||||
static_cast<int32_t>(K_idx),
|
||||
static_cast<int32_t>(L_idx),
|
||||
k_tiles,
|
||||
k_tiles // remaining iterations
|
||||
};
|
||||
}
|
||||
else {
|
||||
// Data-parallel case
|
||||
return {
|
||||
static_cast<int32_t>(work_tile_info.M_idx),
|
||||
static_cast<int32_t>(work_tile_info.N_idx),
|
||||
static_cast<int32_t>(0), // K_idx
|
||||
static_cast<int32_t>(work_tile_info.L_idx),
|
||||
static_cast<uint32_t>(params_.sk_params_.divmod_tiles_per_output_tile_.divisor),
|
||||
static_cast<uint32_t>(params_.sk_params_.divmod_tiles_per_output_tile_.divisor)
|
||||
};
|
||||
}
|
||||
}
|
||||
|
||||
// Converts a WorkTileInfo struct to the WorkTileInfo representation
|
||||
// of the underlying SM100 scheduler.
|
||||
CUTLASS_HOST_DEVICE static
|
||||
InternalWorkTileInfo
|
||||
to_underlying_work_tile_info(WorkTileInfo const& work_tile_info) {
|
||||
return {
|
||||
work_tile_info.M_idx,
|
||||
work_tile_info.N_idx,
|
||||
work_tile_info.L_idx,
|
||||
work_tile_info.is_valid()
|
||||
};
|
||||
}
|
||||
|
||||
// Returns whether the current parameters contain only data-parallel tiles
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
is_dp_only() const {
|
||||
return params_.sk_params_.sk_units_ == 0 && params_.sk_params_.divmod_splits_.divisor == 1;
|
||||
}
|
||||
|
||||
// Returns whether the current parameters are for a split-K decomposition
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
is_split_k() const {
|
||||
return params_.sk_params_.divmod_splits_.divisor > 1;
|
||||
}
|
||||
|
||||
// Returns whether the current parameters contain any stream-K work
|
||||
CUTLASS_HOST_DEVICE
|
||||
bool
|
||||
has_sk_work() const {
|
||||
return params_.sk_params_.sk_units_ > 0;
|
||||
}
|
||||
|
||||
// Performs reduction across splits for a given output tile
|
||||
template <
|
||||
class TiledMma,
|
||||
class AccEngine,
|
||||
class AccLayout,
|
||||
class AccumulatorPipeline,
|
||||
class AccumulatorPipelineState,
|
||||
class CopyOpT2R
|
||||
>
|
||||
CUTLASS_DEVICE
|
||||
AccumulatorPipelineState
|
||||
tmem_fixup(
|
||||
TiledMma const& tiled_mma,
|
||||
WorkTileInfo const& work_tile_info,
|
||||
cute::Tensor<AccEngine, AccLayout>& accumulators,
|
||||
AccumulatorPipeline acc_pipeline,
|
||||
AccumulatorPipelineState acc_pipe_consumer_state,
|
||||
CopyOpT2R,
|
||||
uint32_t num_accumulator_mtx = 1,
|
||||
uint32_t idx_accumulator_mtx = 0) const {
|
||||
using namespace cute;
|
||||
static_assert(cute::is_tmem_v<AccEngine>, "Accumulator must be in TMEM");
|
||||
|
||||
using ElementAccumulator = typename AccEngine::element_type;
|
||||
|
||||
constexpr uint32_t ThreadsForFixup = NumThreadsPerWarpGroup;
|
||||
constexpr uint32_t Offset = static_cast<int>(cutlass::arch::ReservedNamedBarriers::StreamkBarrier0);
|
||||
constexpr uint32_t MaxNumNamedBarriers = 1;
|
||||
constexpr uint32_t BarrierIdx = 0;
|
||||
using BarrierManager = NamedBarrierManager<ThreadsForFixup, Offset, MaxNumNamedBarriers>;
|
||||
|
||||
// When accumulators reside in TMEM, perform TMEM -> RF loads before performing fixup,
|
||||
// and perform RF -> TMEM stores after fixup (when the split must compute the epilogue)
|
||||
auto dummy_gmem_workspace = make_tensor(
|
||||
make_gmem_ptr<ElementAccumulator>(nullptr),
|
||||
make_layout(take<0,2>(TileShape{}), GenRowMajor{})); // (TILE_M,TILE_N)
|
||||
|
||||
auto dummy_gmem_buffer = tiled_mma.get_slice(0).partition_C(dummy_gmem_workspace); // (MMA,MMA_M,MMA_N)
|
||||
|
||||
auto tmem_load = make_tmem_copy(CopyOpT2R{}, accumulators);
|
||||
auto tmem_store = make_tmem_copy(cute::TMEM::tmem_load_to_store(CopyOpT2R{}), accumulators);
|
||||
|
||||
auto thr_tmem_load = tmem_load.get_slice(threadIdx.x % ThreadsForFixup);
|
||||
auto thr_tmem_store = tmem_store.get_slice(threadIdx.x % ThreadsForFixup);
|
||||
|
||||
Tensor tCtAcc = thr_tmem_load.partition_S(accumulators); // (TMEM_LOAD,TMEM_LOAD_MMA,TMEM_LOAD_M,TMEM_LOAD_N)
|
||||
Tensor tCgAcc = thr_tmem_load.partition_D(dummy_gmem_buffer); // (TMEM_LOAD,TMEM_LOAD_MMA,TMEM_LOAD_M,TMEM_LOAD_N)
|
||||
auto tCrAcc = make_tensor<ElementAccumulator>(shape(tCgAcc)); // (TMEM_LOAD,TMEM_LOAD_MMA,TMEM_LOAD_M,TMEM_LOAD_N)
|
||||
|
||||
acc_pipeline.consumer_wait(acc_pipe_consumer_state);
|
||||
|
||||
// Copy accumulators from tmem to rmem for reduction
|
||||
copy(tmem_load, tCtAcc, tCrAcc);
|
||||
|
||||
bool should_compute_epilogue = compute_epilogue(work_tile_info);
|
||||
if (!should_compute_epilogue && (idx_accumulator_mtx == (num_accumulator_mtx - 1))) {
|
||||
// Splits that do not compute the epilogue must advance the accumulator pipeline
|
||||
cutlass::arch::fence_view_async_tmem_load();
|
||||
acc_pipeline.consumer_release(acc_pipe_consumer_state);
|
||||
++acc_pipe_consumer_state;
|
||||
}
|
||||
|
||||
// Perform fixup
|
||||
UnderlyingStreamKScheduler::template fixup_helper<decltype(tCrAcc), BarrierManager>(
|
||||
params_.sk_params_, work_tile_info, tCrAcc, MaxNumNamedBarriers, BarrierIdx, num_accumulator_mtx, idx_accumulator_mtx);
|
||||
|
||||
if (should_compute_epilogue) {
|
||||
// Splits that compute the epilogue copy the reduced accumulators back to tmem for
|
||||
// the epilogue to compute on it
|
||||
copy(tmem_store, tCrAcc, tCtAcc);
|
||||
}
|
||||
|
||||
return acc_pipe_consumer_state;
|
||||
}
|
||||
|
||||
|
||||
//
|
||||
// Members
|
||||
//
|
||||
|
||||
UnderlyingScheduler sm100_scheduler_;
|
||||
Params params_;
|
||||
dim3 block_id_in_cluster_;
|
||||
uint64_t current_work_linear_idx_ = 0;
|
||||
uint32_t unit_iter_start_ = 0;
|
||||
|
||||
// This might not be needed
|
||||
bool is_fallback_cluster_ = false;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // end namespace cutlass::gemm::kernel::detail
|
||||
@@ -120,6 +120,7 @@ public:
|
||||
typename detail::TileSchedulerSelector<
|
||||
GroupScheduler, ArchTag,
|
||||
TileShape, ClusterShape,
|
||||
2, // Default unused parameter - SchedulerPipelineStageCoun
|
||||
ProblemShape>::Scheduler,
|
||||
typename detail::TileSchedulerSelector<
|
||||
void, ArchTag, TileShape, ClusterShape>::Scheduler>;
|
||||
|
||||
@@ -120,6 +120,7 @@ public:
|
||||
typename detail::TileSchedulerSelector<
|
||||
GroupScheduler, ArchTag,
|
||||
TileShape, ClusterShape,
|
||||
2, // Default unused parameter - SchedulerPipelineStageCoun
|
||||
ProblemShape>::Scheduler,
|
||||
typename detail::TileSchedulerSelector<
|
||||
void, ArchTag, TileShape, ClusterShape>::Scheduler>;
|
||||
|
||||
@@ -401,9 +401,7 @@ public:
|
||||
|
||||
// Compute m_coord, n_coord, and l_coord with their post-tiled shapes
|
||||
auto m_coord = idx2crd(int(blockIdx.x), shape<2>(gA_mkl));
|
||||
|
||||
auto n_coord = idx2crd(int(blockIdx.y), shape<2>(gB_nkl), compact_col_major(shape<2>(gB_nkl)));
|
||||
|
||||
auto n_coord = idx2crd(int(blockIdx.y), shape<2>(gB_nkl));
|
||||
// handles the difference between the rank of Tensor returned by load_input in case they do not have a batch mode
|
||||
auto l_coord = [&] (auto const& gB_nkl_) {
|
||||
// gB_nkl needs to be passed into the lambda because C++17
|
||||
|
||||
@@ -714,6 +714,19 @@ public:
|
||||
auto [cta_m_in_cluster_, cta_n_in_cluster_, _] = block_id_in_cluster;
|
||||
uint64_t cta_m_in_cluster = static_cast<uint64_t>(cta_m_in_cluster_);
|
||||
uint64_t cta_n_in_cluster = static_cast<uint64_t>(cta_n_in_cluster_);
|
||||
|
||||
// Determine the CTA's M and N offsets within the preferred cluster
|
||||
// This simply finds the linear offset of the CTA within the cluster, and takes a divmod
|
||||
// on it depending on the rasterization order used by the scheduler.
|
||||
uint64_t cluster_linear_work_idx_tmp = params.div_cluster_size(linear_idx) * params.get_cluster_size();
|
||||
|
||||
if (params.raster_order_ == RasterOrder::AlongN) {
|
||||
params.divmod_cluster_shape_minor_(cta_n_in_cluster, cta_m_in_cluster, linear_idx - cluster_linear_work_idx_tmp);
|
||||
}
|
||||
else {
|
||||
params.divmod_cluster_shape_minor_(cta_m_in_cluster, cta_n_in_cluster, linear_idx - cluster_linear_work_idx_tmp);
|
||||
}
|
||||
|
||||
return {static_cast<uint32_t>(cta_m_in_cluster), static_cast<uint32_t>(cta_n_in_cluster), _};
|
||||
}
|
||||
|
||||
@@ -946,6 +959,8 @@ private:
|
||||
params.divmod_cluster_blk_major_,
|
||||
params.log_swizzle_size_,
|
||||
params.raster_order_
|
||||
, cta_m_in_cluster
|
||||
, cta_n_in_cluster
|
||||
);
|
||||
|
||||
// Set the M, N, and L block offsets
|
||||
|
||||
@@ -58,6 +58,9 @@ struct GroupScheduler { }; // Only used for Grouped GEMMs
|
||||
#include "cutlass/gemm/kernel/sm90_tile_scheduler.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_tile_scheduler_stream_k.hpp"
|
||||
#include "cutlass/gemm/kernel/sm90_tile_scheduler_group.hpp"
|
||||
#include "cutlass/gemm/kernel/sm100_tile_scheduler.hpp"
|
||||
#include "cutlass/gemm/kernel/sm100_tile_scheduler_stream_k.hpp"
|
||||
#include "cutlass/gemm/kernel/sm100_tile_scheduler_group.hpp"
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::kernel::detail {
|
||||
@@ -71,6 +74,7 @@ template <
|
||||
class ArchTag,
|
||||
class TileShape,
|
||||
class ClusterShape
|
||||
, uint32_t SchedulerPipelineStageCount = 2
|
||||
, class ProblemShapeType = void
|
||||
>
|
||||
struct TileSchedulerSelector {
|
||||
@@ -82,12 +86,14 @@ template <
|
||||
class ArchTag,
|
||||
class TileShape,
|
||||
class ClusterShape
|
||||
, uint32_t SchedulerPipelineStageCount
|
||||
>
|
||||
struct TileSchedulerSelector<
|
||||
PersistentScheduler,
|
||||
ArchTag,
|
||||
TileShape,
|
||||
ClusterShape
|
||||
, SchedulerPipelineStageCount
|
||||
> {
|
||||
using Scheduler = PersistentTileSchedulerSm90;
|
||||
};
|
||||
@@ -97,30 +103,35 @@ template <
|
||||
class ArchTag,
|
||||
class TileShape,
|
||||
class ClusterShape
|
||||
, uint32_t SchedulerPipelineStageCount
|
||||
>
|
||||
struct TileSchedulerSelector<
|
||||
void,
|
||||
ArchTag,
|
||||
TileShape,
|
||||
ClusterShape
|
||||
, SchedulerPipelineStageCount
|
||||
> {
|
||||
using Scheduler = typename TileSchedulerSelector<
|
||||
PersistentScheduler,
|
||||
ArchTag,
|
||||
TileShape,
|
||||
ClusterShape
|
||||
, SchedulerPipelineStageCount
|
||||
>::Scheduler;
|
||||
};
|
||||
|
||||
template <
|
||||
class TileShape,
|
||||
class ClusterShape
|
||||
, uint32_t SchedulerPipelineStageCount
|
||||
>
|
||||
struct TileSchedulerSelector<
|
||||
StreamKScheduler,
|
||||
arch::Sm90,
|
||||
TileShape,
|
||||
ClusterShape
|
||||
, SchedulerPipelineStageCount
|
||||
> {
|
||||
using Scheduler = PersistentTileSchedulerSm90StreamK<TileShape, ClusterShape>;
|
||||
};
|
||||
@@ -128,6 +139,7 @@ struct TileSchedulerSelector<
|
||||
template <
|
||||
class TileShape,
|
||||
class ClusterShape
|
||||
, uint32_t SchedulerPipelineStageCount
|
||||
, class GroupProblemShape
|
||||
>
|
||||
struct TileSchedulerSelector<
|
||||
@@ -135,11 +147,113 @@ struct TileSchedulerSelector<
|
||||
arch::Sm90,
|
||||
TileShape,
|
||||
ClusterShape
|
||||
, SchedulerPipelineStageCount
|
||||
, GroupProblemShape
|
||||
> {
|
||||
using Scheduler = PersistentTileSchedulerSm90Group<GroupProblemShape>;
|
||||
};
|
||||
|
||||
|
||||
template <class TileShape, class ClusterShape, uint32_t SchedulerPipelineStageCount>
|
||||
struct TileSchedulerSelector<
|
||||
PersistentScheduler,
|
||||
arch::Sm100,
|
||||
TileShape,
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount> {
|
||||
using Scheduler = PersistentTileSchedulerSm100<
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount>;
|
||||
};
|
||||
|
||||
// Ptr-Array kernel may provide a specialized ArrayProblemShape type
|
||||
template <class TileShape,
|
||||
class ClusterShape,
|
||||
uint32_t SchedulerPipelineStageCount,
|
||||
class ProblemShape>
|
||||
struct TileSchedulerSelector<
|
||||
PersistentScheduler,
|
||||
arch::Sm100,
|
||||
TileShape,
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount,
|
||||
ProblemShape> {
|
||||
using Scheduler = PersistentTileSchedulerSm100<
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount>;
|
||||
};
|
||||
|
||||
// Default (void) for Sm100 maps to PersistentTileSchedulerSm100
|
||||
template <class TileShape, class ClusterShape, uint32_t SchedulerPipelineStageCount>
|
||||
struct TileSchedulerSelector<
|
||||
void,
|
||||
arch::Sm100,
|
||||
TileShape,
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount> {
|
||||
using Scheduler = typename TileSchedulerSelector<
|
||||
PersistentScheduler,
|
||||
arch::Sm100,
|
||||
TileShape,
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount>::Scheduler;
|
||||
};
|
||||
|
||||
// Default (void) for Sm100 maps to PersistentTileSchedulerSm100
|
||||
// Ptr-Array kernel may provide a specialized ArrayProblemShape type
|
||||
template <class TileShape,
|
||||
class ClusterShape,
|
||||
uint32_t SchedulerPipelineStageCount,
|
||||
class ProblemShape>
|
||||
struct TileSchedulerSelector<
|
||||
void,
|
||||
arch::Sm100,
|
||||
TileShape,
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount,
|
||||
ProblemShape> {
|
||||
using Scheduler = typename TileSchedulerSelector<
|
||||
PersistentScheduler,
|
||||
arch::Sm100,
|
||||
TileShape,
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount>::Scheduler;
|
||||
};
|
||||
|
||||
// SM100 Group tile scheduler
|
||||
template <
|
||||
class TileShape,
|
||||
class ClusterShape,
|
||||
uint32_t SchedulerPipelineStageCount,
|
||||
class GroupProblemShape
|
||||
>
|
||||
struct TileSchedulerSelector<
|
||||
GroupScheduler,
|
||||
arch::Sm100,
|
||||
TileShape,
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount,
|
||||
GroupProblemShape
|
||||
> {
|
||||
using Scheduler = PersistentTileSchedulerSm100Group<GroupProblemShape>;
|
||||
};
|
||||
|
||||
// SM100 stream-K scheduler
|
||||
template <class TileShape, class ClusterShape, uint32_t SchedulerPipelineStageCount>
|
||||
struct TileSchedulerSelector<
|
||||
StreamKScheduler,
|
||||
arch::Sm100,
|
||||
TileShape,
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount> {
|
||||
using Scheduler = PersistentTileSchedulerSm100StreamK<
|
||||
TileShape,
|
||||
ClusterShape,
|
||||
SchedulerPipelineStageCount>;
|
||||
};
|
||||
|
||||
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm::kernel::detail
|
||||
|
||||
@@ -189,6 +189,7 @@ struct PersistentTileSchedulerSm90Params {
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option,
|
||||
bool truncate_by_problem_size=true
|
||||
, bool bypass_occupancy_calculation=false
|
||||
) {
|
||||
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape, cta_shape, cluster_shape);
|
||||
@@ -199,6 +200,7 @@ struct PersistentTileSchedulerSm90Params {
|
||||
max_swizzle_size,
|
||||
raster_order_option,
|
||||
truncate_by_problem_size
|
||||
, bypass_occupancy_calculation
|
||||
);
|
||||
}
|
||||
|
||||
@@ -214,6 +216,7 @@ struct PersistentTileSchedulerSm90Params {
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option,
|
||||
bool truncate_by_problem_size=true
|
||||
, bool bypass_occupancy_calculation=false
|
||||
) {
|
||||
|
||||
int const sm_count = hw_info.sm_count;
|
||||
@@ -274,6 +277,7 @@ struct PersistentTileSchedulerSm90Params {
|
||||
}
|
||||
else {
|
||||
int cta_per_device = sm_count;
|
||||
if (!bypass_occupancy_calculation) {
|
||||
/*
|
||||
* Optimal grid size calculation is based on
|
||||
* GH100: 8 GPCs, 72 TPCs (9 TPCs/GPC), 2 SMs/TPC, 144 SMs per full GPU
|
||||
@@ -281,6 +285,8 @@ struct PersistentTileSchedulerSm90Params {
|
||||
*/
|
||||
constexpr int max_sm_per_gpc = 18;
|
||||
cta_per_device = get_max_cta_occupancy(max_sm_per_gpc, cluster_shape, sm_count);
|
||||
}
|
||||
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid.y = possibly_truncate(
|
||||
cta_per_device / cluster_shape.m(),
|
||||
@@ -380,7 +386,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
// Strategies for computing reductions between CTAs computing portions of a given output tile
|
||||
enum class ReductionMode {
|
||||
// Participating CTAs perform reduction in a turnstile fashion in order of the K extent
|
||||
// covered by each CTA. This requires a lock to be held exclusively be the CTA that is
|
||||
// covered by each CTA. This requires a lock to be held exclusively by the CTA that is
|
||||
// currently accumulating.
|
||||
//
|
||||
// Turnstile accumulation ensures deterministic numeric behavior when using this mode.
|
||||
@@ -502,6 +508,32 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
// Divides dividend by the cluster size in the M dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
uint64_t
|
||||
truncate_to_cluster_size_m(uint64_t dividend) const {
|
||||
if (raster_order_ == RasterOrder::AlongN) {
|
||||
return divmod_cluster_shape_minor_.divide(dividend) * divmod_cluster_shape_minor_.divisor;
|
||||
}
|
||||
else {
|
||||
return divmod_cluster_shape_major_.divide(dividend) * divmod_cluster_shape_major_.divisor;
|
||||
}
|
||||
}
|
||||
|
||||
// Divides dividend by the cluster size in the N dimension
|
||||
CUTLASS_HOST_DEVICE
|
||||
uint64_t
|
||||
truncate_to_cluster_size_n(uint64_t dividend) const {
|
||||
if (raster_order_ == RasterOrder::AlongM) {
|
||||
return divmod_cluster_shape_minor_.divide(dividend) * divmod_cluster_shape_minor_.divisor;
|
||||
}
|
||||
else {
|
||||
return divmod_cluster_shape_major_.divide(dividend) * divmod_cluster_shape_major_.divisor;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
uint64_t
|
||||
get_cluster_size() const {
|
||||
@@ -542,6 +574,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
DecompositionMode decomposition_mode,
|
||||
void* workspace,
|
||||
const uint32_t epilogue_subtile = 1u
|
||||
, uint32_t ktile_start_alignment_count = 1u
|
||||
) {
|
||||
dim3 problem_blocks = UnderlyingParams::get_tiled_cta_shape_mnl(
|
||||
problem_shape, tile_shape, cluster_shape);
|
||||
@@ -561,6 +594,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
decomposition_mode,
|
||||
workspace,
|
||||
epilogue_subtile
|
||||
, ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
@@ -580,6 +614,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
DecompositionMode decomposition_mode,
|
||||
void* workspace,
|
||||
const uint32_t epilogue_subtile = 1
|
||||
, uint32_t ktile_start_alignment_count = 1u
|
||||
) {
|
||||
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
@@ -590,6 +625,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
}
|
||||
#endif // !defined(__CUDACC_RTC__)
|
||||
|
||||
ktile_start_alignment_count_ = ktile_start_alignment_count;
|
||||
UnderlyingParams underlying_params;
|
||||
underlying_params.initialize(
|
||||
problem_blocks,
|
||||
@@ -716,6 +752,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
DecompositionMode decomposition_mode,
|
||||
ReductionMode reduction_mode,
|
||||
const uint32_t epilogue_subtile = 1
|
||||
, uint32_t ktile_start_alignment_count = 1u
|
||||
) {
|
||||
uint32_t groups = 0;
|
||||
uint32_t sk_tiles = 0;
|
||||
@@ -749,6 +786,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
decomposition_mode,
|
||||
reduction_mode,
|
||||
epilogue_subtile
|
||||
, ktile_start_alignment_count
|
||||
);
|
||||
|
||||
// Given heuristic_mode returned from the heuristic() method, set params fields.
|
||||
@@ -772,6 +810,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
splits,
|
||||
epilogue_subtile,
|
||||
reduction_mode
|
||||
, ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
@@ -797,6 +836,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
DecompositionMode decomposition_mode,
|
||||
ReductionMode reduction_mode,
|
||||
uint32_t epilogue_subtile
|
||||
, uint32_t ktile_start_alignment_count
|
||||
) {
|
||||
|
||||
// Get block numbers in m, n and l dimensions
|
||||
@@ -805,6 +845,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
// Short circuit to basic split-K decomposition
|
||||
uint32_t adapted_splits = adjust_split_count(
|
||||
splits, hw_info.sm_count, k_tiles_per_output_tile
|
||||
, ktile_start_alignment_count
|
||||
);
|
||||
sk_splits = adapted_splits;
|
||||
return DecompositionMode::SplitK;
|
||||
@@ -826,6 +867,8 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
);
|
||||
uint64_t ctas_per_wave = grid.x * grid.y;
|
||||
cluster_size = cluster_shape.m() * cluster_shape.n();
|
||||
uint64_t ctas_per_wave_in_full_clusters = (ctas_per_wave / cluster_size) * cluster_size;
|
||||
|
||||
// The number of output tiles to be computed in stream-K and data-parallel fashion, respectively.
|
||||
sk_tiles = get_num_sk_tiles(
|
||||
output_tiles,
|
||||
@@ -833,6 +876,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
cluster_size,
|
||||
k_tiles_per_output_tile,
|
||||
decomposition_mode
|
||||
, ctas_per_wave_in_full_clusters
|
||||
);
|
||||
uint64_t dp_tiles = output_tiles - sk_tiles;
|
||||
// Calculate the number of work units covering the data-parallel and stream-K tiles.
|
||||
@@ -846,6 +890,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
dp_units = dp_tiles;
|
||||
|
||||
uint64_t ctas_per_sk_wave = ctas_per_wave;
|
||||
ctas_per_sk_wave = ctas_per_wave_in_full_clusters;
|
||||
sk_units = get_num_sk_units(cluster_shape, ctas_per_sk_wave, sk_tiles, k_tiles_per_output_tile);
|
||||
|
||||
if (decomposition_mode == DecompositionMode::DataParallel ||
|
||||
@@ -924,6 +969,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
uint32_t splits,
|
||||
uint32_t epilogue_subtile,
|
||||
ReductionMode reduction_mode
|
||||
, uint32_t ktile_start_alignment_count
|
||||
) {
|
||||
// The highest priority when customers set as splitk mode, may set
|
||||
// with a adpated splits value rather than the original splits
|
||||
@@ -1025,6 +1071,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
max_swizzle_size,
|
||||
raster_order_option,
|
||||
/* truncate_by_problem_size = */false
|
||||
/* bypass_occupancy_calculation = */, true
|
||||
);
|
||||
}
|
||||
|
||||
@@ -1037,6 +1084,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
uint64_t cluster_size,
|
||||
uint32_t k_tiles_per_output_tile,
|
||||
DecompositionMode decomposition_mode
|
||||
, uint64_t ctas_per_wave_in_full_clusters
|
||||
) {
|
||||
uint32_t full_waves = static_cast<uint32_t>(output_tiles / ctas_per_wave);
|
||||
uint32_t total_waves = static_cast<uint32_t>((output_tiles + ctas_per_wave - 1) / ctas_per_wave);
|
||||
@@ -1054,16 +1102,16 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
uint64_t dp_tiles = dp_waves * ctas_per_wave;
|
||||
uint64_t sk_tiles = output_tiles - dp_tiles;
|
||||
|
||||
if (decomposition_mode == DecompositionMode::Heuristic) {
|
||||
if (full_waves == total_waves || k_tiles_per_output_tile <= min_iters_per_sk_unit_) {
|
||||
// All tiles will be data-parallel tiles if there is either no quantization
|
||||
// or if there is no work to be split.
|
||||
return 0;
|
||||
}
|
||||
//
|
||||
// The final wave is not full. Perform some stream-K work.
|
||||
//
|
||||
if (full_waves == total_waves || k_tiles_per_output_tile <= min_iters_per_sk_unit_) {
|
||||
// All tiles will be data-parallel tiles if there is either no quantization
|
||||
// or if there is no work to be split.
|
||||
return 0;
|
||||
}
|
||||
|
||||
//
|
||||
// The final wave is not full. Perform some stream-K work.
|
||||
//
|
||||
if (decomposition_mode == DecompositionMode::Heuristic) {
|
||||
// Rudimentary heuristic: prefer data-parallel decomposition if we have more than
|
||||
// one wave and the tail wave is more than half full. This is subject to change.
|
||||
uint64_t tail_tiles = output_tiles - (full_waves * ctas_per_wave);
|
||||
@@ -1071,7 +1119,6 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
return 0;
|
||||
}
|
||||
}
|
||||
|
||||
return static_cast<uint32_t>(sk_tiles);
|
||||
}
|
||||
|
||||
@@ -1172,14 +1219,17 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
);
|
||||
uint64_t ctas_per_wave = grid.x * grid.y;
|
||||
uint64_t cluster_size = cluster_shape.m() * cluster_shape.n();
|
||||
uint64_t ctas_per_wave_in_full_clusters = (ctas_per_wave / cluster_size) * cluster_size;
|
||||
uint32_t sk_tiles = get_num_sk_tiles(
|
||||
output_tiles,
|
||||
ctas_per_wave,
|
||||
cluster_size,
|
||||
static_cast<uint32_t>(k_tiles_per_output_tile),
|
||||
decomposition_mode
|
||||
, ctas_per_wave_in_full_clusters
|
||||
);
|
||||
uint64_t ctas_per_sk_wave = ctas_per_wave;
|
||||
ctas_per_sk_wave = ctas_per_wave_in_full_clusters;
|
||||
uint64_t sk_units = get_num_sk_units(cluster_shape, ctas_per_sk_wave, sk_tiles, k_tiles_per_output_tile);
|
||||
uint64_t dp_tiles = output_tiles - sk_tiles;
|
||||
|
||||
@@ -1187,11 +1237,13 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
(decomposition_mode == DecompositionMode::Heuristic && splits > 1)) {
|
||||
splits = adjust_split_count(
|
||||
splits, new_hw_info.sm_count, k_tiles_per_output_tile
|
||||
, ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
bool split_k_required = splits > 1 && (decomposition_mode == DecompositionMode::SplitK || decomposition_mode == DecompositionMode::Heuristic);
|
||||
bool split_k_selected = decomposition_mode == DecompositionMode::Heuristic &&
|
||||
bool split_k_selected = !split_k_required &&
|
||||
decomposition_mode == DecompositionMode::Heuristic &&
|
||||
sk_units > sk_tiles &&
|
||||
sk_tiles != 0 &&
|
||||
sk_units % sk_tiles == 0;
|
||||
@@ -1547,6 +1599,7 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
int splits,
|
||||
int sm_count,
|
||||
uint32_t k_tiles_per_output_tile
|
||||
, uint32_t ktile_start_alignment_count
|
||||
) {
|
||||
// Don't split by more than the available number of SMs
|
||||
if (splits > sm_count) {
|
||||
@@ -1561,6 +1614,11 @@ struct PersistentTileSchedulerSm90StreamKParams {
|
||||
// If k_tiles_per_output_tiles / splits == 1, there will be one k_tile per cta
|
||||
// and this violate k_tile start from even requirements. Thus we need to
|
||||
// reduce the number of splits.
|
||||
if (ktile_start_alignment_count > 1u &&
|
||||
splits > 1 &&
|
||||
k_tiles_per_output_tile / static_cast<uint32_t>(splits) == 1) {
|
||||
splits = k_tiles_per_output_tile / ktile_start_alignment_count;
|
||||
}
|
||||
return splits;
|
||||
}
|
||||
};
|
||||
@@ -1809,6 +1867,732 @@ struct PersistentTileSchedulerSm90GroupParams {
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
//
|
||||
// Parameters for SM100 tile schedulers
|
||||
//
|
||||
|
||||
// Parameters for SM100 persistent tile scheduler
|
||||
struct PersistentTileSchedulerSm100Params {
|
||||
|
||||
using UnderlyingParams = PersistentTileSchedulerSm90Params;
|
||||
|
||||
using RasterOrder = UnderlyingParams::RasterOrder;
|
||||
using RasterOrderOptions = UnderlyingParams::RasterOrderOptions;
|
||||
|
||||
uint32_t problem_tiles_m_ = 0;
|
||||
uint32_t problem_tiles_n_ = 0;
|
||||
uint32_t problem_tiles_l_ = 0;
|
||||
FastDivmod divmod_cluster_shape_m_{};
|
||||
FastDivmod divmod_cluster_shape_n_{};
|
||||
RasterOrder raster_order_ = RasterOrder::AlongM;
|
||||
int32_t log_swizzle_size_ = 0;
|
||||
|
||||
// Initializes members. This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
void
|
||||
initialize(
|
||||
BatchedGemmCoord problem_shape,
|
||||
GemmCoord tile_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option
|
||||
) {
|
||||
dim3 problem_blocks = UnderlyingParams::get_tiled_cta_shape_mnl(problem_shape, tile_shape, cluster_shape);
|
||||
initialize(
|
||||
problem_blocks,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
max_swizzle_size,
|
||||
raster_order_option
|
||||
);
|
||||
}
|
||||
|
||||
// Version of initialize that takes in as input the number of CTAs in the M and N and L dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
void
|
||||
initialize(
|
||||
dim3 problem_blocks,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option
|
||||
) {
|
||||
|
||||
CUTLASS_UNUSED(hw_info);
|
||||
CUTLASS_UNUSED(max_swizzle_size);
|
||||
|
||||
// Cluster counters in m, n and l dimensions of the problem tiles
|
||||
problem_tiles_m_ = problem_blocks.x / cluster_shape.m();
|
||||
problem_tiles_n_ = problem_blocks.y / cluster_shape.n();
|
||||
problem_tiles_l_ = problem_blocks.z;
|
||||
divmod_cluster_shape_m_ = FastDivmod(cluster_shape.m());
|
||||
divmod_cluster_shape_n_ = FastDivmod(cluster_shape.n());
|
||||
|
||||
raster_order_ = UnderlyingParams::get_rasterization_order(problem_tiles_m_, problem_tiles_n_, raster_order_option);
|
||||
if (raster_order_option == RasterOrderOptions::Heuristic && raster_order_ == RasterOrder::AlongN) {
|
||||
// The current implementation of AlongN rasterization for B100 requires swapping the number of clusters along the
|
||||
// X and Y dimensions of the grid. However, since the grid Y dimension has a smaller range of allowed values
|
||||
// than the grid X dimension, we must check whether the swapped grid would exceed the grid Y limit. If the
|
||||
// swapped grid would exceed this limit, simply rever to AlongM mode.
|
||||
//
|
||||
// Overflow in the swapped X dimension is not possible. At worst, there will be ((1 << 16) - 1) clusters
|
||||
// along the original Y dimension of the grid. Even if the cluster M mode is 16, the new grid X value
|
||||
// will be at most ((1 << 16) - 1) * 16, which is less than the grid X limit of ((1 << 31) - 1).
|
||||
uint32_t cluster_m = static_cast<uint32_t>(problem_blocks.x) / static_cast<uint32_t>(cluster_shape.m());
|
||||
uint32_t new_grid_y = cluster_m * static_cast<uint32_t>(cluster_shape.n());
|
||||
|
||||
if (new_grid_y > (1 << 16) - 1) {
|
||||
raster_order_ = RasterOrder::AlongM;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Given the inputs, computes the physical grid we should launch.
|
||||
// This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
CUTLASS_HOST_DEVICE static
|
||||
dim3
|
||||
get_grid_shape(
|
||||
BatchedGemmCoord problem_shape,
|
||||
GemmCoord cta_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo hw_info,
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option
|
||||
) {
|
||||
|
||||
CUTLASS_UNUSED(cluster_shape);
|
||||
CUTLASS_UNUSED(hw_info);
|
||||
CUTLASS_UNUSED(max_swizzle_size);
|
||||
CUTLASS_UNUSED(raster_order_option);
|
||||
|
||||
return get_tiled_cta_shape_mnl(problem_shape, cta_shape, cluster_shape);
|
||||
}
|
||||
|
||||
// Get the number of CTA tiles in this problem. This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_tiled_cta_shape_mnl(
|
||||
BatchedGemmCoord problem_shape,
|
||||
GemmCoord cta_shape,
|
||||
GemmCoord cluster_shape) {
|
||||
|
||||
return UnderlyingParams::get_tiled_cta_shape_mnl(problem_shape, cta_shape, cluster_shape);
|
||||
}
|
||||
|
||||
// Get the amount of scratch workspace needed for the kernel. This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
static size_t
|
||||
get_workspace_size(
|
||||
BatchedGemmCoord problem_shape,
|
||||
GemmCoord tile_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option
|
||||
) {
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape, tile_shape, cluster_shape);
|
||||
return get_workspace_size(
|
||||
problem_blocks,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
max_swizzle,
|
||||
raster_order_option
|
||||
);
|
||||
}
|
||||
|
||||
// Version of get_workspace_size that takes in as input the number of CTAs in the M and N dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
static size_t
|
||||
get_workspace_size(
|
||||
dim3 problem_blocks,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option
|
||||
) {
|
||||
|
||||
CUTLASS_UNUSED(problem_blocks);
|
||||
CUTLASS_UNUSED(cluster_shape);
|
||||
CUTLASS_UNUSED(hw_info);
|
||||
CUTLASS_UNUSED(max_swizzle);
|
||||
CUTLASS_UNUSED(raster_order_option);
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
// Initialize the workspace to be used for the kernel. This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
static cutlass::Status
|
||||
initialize_workspace(
|
||||
void* workspace,
|
||||
cudaStream_t stream,
|
||||
BatchedGemmCoord problem_shape,
|
||||
GemmCoord tile_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
CudaHostAdapter *cuda_adapter = nullptr
|
||||
) {
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape, tile_shape, cluster_shape);
|
||||
return initialize_workspace(
|
||||
workspace,
|
||||
stream,
|
||||
problem_blocks,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
max_swizzle,
|
||||
raster_order_option,
|
||||
cuda_adapter
|
||||
);
|
||||
}
|
||||
|
||||
// Version of initialize_workspace that takes in as input the number of CTAs in the M and N dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
static cutlass::Status
|
||||
initialize_workspace(
|
||||
void* workspace,
|
||||
cudaStream_t stream,
|
||||
dim3 problem_blocks,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
CudaHostAdapter *cuda_adapter = nullptr
|
||||
) {
|
||||
|
||||
CUTLASS_UNUSED(workspace);
|
||||
CUTLASS_UNUSED(stream);
|
||||
CUTLASS_UNUSED(problem_blocks);
|
||||
CUTLASS_UNUSED(cluster_shape);
|
||||
CUTLASS_UNUSED(hw_info);
|
||||
CUTLASS_UNUSED(max_swizzle);
|
||||
CUTLASS_UNUSED(raster_order_option);
|
||||
|
||||
return cutlass::Status::kSuccess;
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Parameters for SM100 persistent stream-K tile scheduler
|
||||
struct PersistentTileSchedulerSm100StreamKParams {
|
||||
|
||||
using UnderlyingParams = PersistentTileSchedulerSm100Params;
|
||||
using UnderlyingStreamKParams = PersistentTileSchedulerSm90StreamKParams;
|
||||
using RasterOrderOptions = UnderlyingParams::RasterOrderOptions;
|
||||
using ReductionMode = UnderlyingStreamKParams::ReductionMode;
|
||||
using DecompositionMode = UnderlyingStreamKParams::DecompositionMode;
|
||||
|
||||
using RasterOrder = UnderlyingParams::RasterOrder;
|
||||
RasterOrder raster_order_ = RasterOrder::AlongM;
|
||||
int32_t log_swizzle_size_ = 0;
|
||||
|
||||
UnderlyingStreamKParams sk_params_{};
|
||||
UnderlyingParams sm100_params_{};
|
||||
|
||||
// Initializes members. This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
void
|
||||
initialize(
|
||||
BatchedGemmCoord problem_shape,
|
||||
GemmCoord tile_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int splits,
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option,
|
||||
ReductionMode reduction_mode,
|
||||
DecompositionMode decomposition_mode,
|
||||
void* workspace,
|
||||
uint32_t ktile_start_alignment_count = 1u
|
||||
) {
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape, tile_shape, cluster_shape);
|
||||
|
||||
// Number of k tiles in each output tile
|
||||
uint32_t k_tiles_per_output_tile = (problem_shape.k() + tile_shape.k() - 1) / tile_shape.k();
|
||||
|
||||
initialize(
|
||||
problem_blocks,
|
||||
k_tiles_per_output_tile,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
splits,
|
||||
max_swizzle_size,
|
||||
raster_order_option,
|
||||
reduction_mode,
|
||||
decomposition_mode,
|
||||
workspace,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
// Version of initialize that takes in as input the number of CTAs in the M and N and L dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
void
|
||||
initialize(
|
||||
dim3 problem_blocks,
|
||||
uint32_t k_tile_per_output_tile,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int splits,
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option,
|
||||
ReductionMode reduction_mode,
|
||||
DecompositionMode decomposition_mode,
|
||||
void* workspace,
|
||||
uint32_t ktile_start_alignment_count = 1u
|
||||
) {
|
||||
sk_params_.initialize(
|
||||
problem_blocks,
|
||||
k_tile_per_output_tile,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
splits,
|
||||
max_swizzle_size,
|
||||
raster_order_option,
|
||||
reduction_mode,
|
||||
decomposition_mode,
|
||||
workspace,
|
||||
/*epilogue_subtile=*/1,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
|
||||
log_swizzle_size_ = sk_params_.log_swizzle_size_;
|
||||
raster_order_ = sk_params_.raster_order_;
|
||||
|
||||
sm100_params_.initialize(
|
||||
problem_blocks,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
max_swizzle_size,
|
||||
RasterOrderOptions::AlongM // Override raster_order to be AlongM, since the SM100 stream-K scheduler does not require grid swapping for raster order selection
|
||||
);
|
||||
}
|
||||
|
||||
// Get the number of CTA tiles in this problem.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_tiled_cta_shape_mnl(
|
||||
BatchedGemmCoord problem_shape,
|
||||
GemmCoord cta_shape,
|
||||
GemmCoord cluster_shape) {
|
||||
|
||||
return UnderlyingParams::get_tiled_cta_shape_mnl(problem_shape, cta_shape, cluster_shape);
|
||||
}
|
||||
|
||||
// Given the inputs, computes the physical grid we should launch.
|
||||
// This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
CUTLASS_HOST_DEVICE
|
||||
dim3
|
||||
get_grid_shape(BatchedGemmCoord problem_shape, GemmCoord cta_shape, GemmCoord cluster_shape) const {
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape, cta_shape, cluster_shape);
|
||||
|
||||
return get_grid_shape(problem_blocks, cluster_shape);
|
||||
}
|
||||
|
||||
// Version of get_grid_shape that takes in as input the number of CTAs in the M and N and L dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
CUTLASS_HOST_DEVICE
|
||||
dim3
|
||||
get_grid_shape(dim3 problem_blocks, GemmCoord cluster_shape) const {
|
||||
if (sk_params_.sk_units_ > 0) {
|
||||
// For stream-K cases, we would, ideally, launch a linear grid of size `sk_params_.units_per_problem_`.
|
||||
// However doing so raises two potential issues:
|
||||
// (a) the total number of tiles in the kernel may exceed the amount that can fit in a single
|
||||
// returned value of a CLC query
|
||||
// (b) the launched grid would not respect cluster-size divisibility requirements
|
||||
//
|
||||
// To circumvent these issues, we must distribute the `sk_params_.units_per_problem_` units of work
|
||||
// across the X, Y, and Z dimensions of the grid, while ensuring that the X and Y dimensions are
|
||||
// divisible by cluster size (we ignore Z, as all CUTLASS kernels currently use a cluster shape
|
||||
// of 1 in the Z dimension).
|
||||
//
|
||||
// For convenience, we launch this as "waves" of `sk_params_.sk_units_` CTAs, with the wave count being
|
||||
// the Z dimension of the grid, and the `sk_params_.sk_units_` CTAs per wave being distributed across
|
||||
// the X and Y dimensions of the grid in a way that alingns with cluster divisibility requirements.
|
||||
//
|
||||
// Thus, the grid that is launched looks like:
|
||||
// grid = dim3(sk_units_ / cluster.y, cluster.y, waves)
|
||||
//
|
||||
// We place sk_units_ / cluster.y in the X dimension of the grid because the CLC query feature
|
||||
// allocates more bits for the X index values returned in the query.
|
||||
//
|
||||
|
||||
// For most cases, `sk_params_.sk_units_` will equal the number of available SMs, so this grid will
|
||||
// naturally represent waves in the true hardware sense.
|
||||
//
|
||||
// However, there are some corner cases in which fewer stream-K units are used than the full SM count
|
||||
// (e.g., if using the full SM count would result in stream-K units that are assigned fewer than the
|
||||
// minimum number of K tile iterations). In these cases, `sk_params_.units_per_problem_` may not be
|
||||
// divisible by `sk_params_.sk_units_`, since any data-parallel work performed alongside stream-K
|
||||
// work is always done in terms of waves of CTAs of number equal to the number of available SMs.
|
||||
// Therefore, we take the ceiling of the division when determining wave count, and allow the underlying
|
||||
// stream-K scheduler to determine which indices are in bounds.
|
||||
uint32_t waves = static_cast<uint32_t>(
|
||||
(sk_params_.units_per_problem_ + sk_params_.sk_units_ - 1) / sk_params_.sk_units_);
|
||||
|
||||
return dim3(
|
||||
sk_params_.sk_units_ / cluster_shape.n(),
|
||||
cluster_shape.n(),
|
||||
waves
|
||||
);
|
||||
}
|
||||
else {
|
||||
// Grid launch for data-parallel and basic split-K decomposition. When data-parallel
|
||||
// mode is used, params.sk_params_.splits = 1.
|
||||
return dim3(problem_blocks.x, problem_blocks.y, problem_blocks.z * sk_params_.divmod_splits_.divisor);
|
||||
}
|
||||
}
|
||||
|
||||
// Get the amount of scratch workspace needed for the kernel. This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
static size_t
|
||||
get_workspace_size(
|
||||
BatchedGemmCoord problem_shape,
|
||||
GemmCoord tile_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int splits,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
DecompositionMode decomposition_mode,
|
||||
ReductionMode reduction_mode,
|
||||
uint32_t reduction_warp_groups,
|
||||
uint32_t barrier_bits,
|
||||
uint32_t element_accumulator_bits,
|
||||
uint32_t ktile_start_alignment_count = 1
|
||||
) {
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape, tile_shape, cluster_shape);
|
||||
uint32_t k_tiles_per_output_tile = (problem_shape.k() + tile_shape.k() - 1) / tile_shape.k();
|
||||
|
||||
return get_workspace_size(
|
||||
problem_blocks,
|
||||
k_tiles_per_output_tile,
|
||||
tile_shape,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
splits,
|
||||
max_swizzle,
|
||||
raster_order_option,
|
||||
decomposition_mode,
|
||||
reduction_mode,
|
||||
reduction_warp_groups,
|
||||
barrier_bits,
|
||||
element_accumulator_bits,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
// Version of get_workspace_size that takes in as input the number of CTAs in the M and N dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
static size_t
|
||||
get_workspace_size(
|
||||
dim3 problem_blocks,
|
||||
uint32_t k_tiles_per_output_tile,
|
||||
GemmCoord tile_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int splits,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
DecompositionMode decomposition_mode,
|
||||
ReductionMode reduction_mode,
|
||||
uint32_t reduction_warp_groups,
|
||||
uint32_t barrier_bits,
|
||||
uint32_t element_accumulator_bits,
|
||||
uint32_t epilogue_subtile = 1,
|
||||
uint32_t num_accumulator_mtxs = 1,
|
||||
uint32_t ktile_start_alignment_count = 1
|
||||
) {
|
||||
return UnderlyingStreamKParams::get_workspace_size(
|
||||
problem_blocks,
|
||||
k_tiles_per_output_tile,
|
||||
tile_shape,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
splits,
|
||||
max_swizzle,
|
||||
raster_order_option,
|
||||
decomposition_mode,
|
||||
reduction_mode,
|
||||
reduction_warp_groups,
|
||||
barrier_bits,
|
||||
element_accumulator_bits,
|
||||
epilogue_subtile,
|
||||
num_accumulator_mtxs,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
// Initialize the workspace to be used for the kernel. This variant of the method should only be used when
|
||||
// problem_shape and tile_shape contain modes of only rank 1.
|
||||
static cutlass::Status
|
||||
initialize_workspace(
|
||||
void* workspace,
|
||||
cudaStream_t stream,
|
||||
BatchedGemmCoord problem_shape,
|
||||
GemmCoord tile_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int splits,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
DecompositionMode decomposition_mode,
|
||||
ReductionMode reduction_mode,
|
||||
uint32_t reduction_warp_groups,
|
||||
uint32_t barrier_bits,
|
||||
uint32_t element_accumulator_bits,
|
||||
uint32_t epilogue_subtile = 1,
|
||||
uint32_t num_accumulator_mtxs = 1,
|
||||
CudaHostAdapter *cuda_adapter = nullptr,
|
||||
uint32_t ktile_start_alignment_count = 1
|
||||
) {
|
||||
dim3 problem_blocks = get_tiled_cta_shape_mnl(problem_shape, tile_shape, cluster_shape);
|
||||
uint32_t k_tiles_per_output_tile = (problem_shape.k() + tile_shape.k() - 1) / tile_shape.k();
|
||||
|
||||
return initialize_workspace(
|
||||
workspace,
|
||||
stream,
|
||||
problem_blocks,
|
||||
k_tiles_per_output_tile,
|
||||
tile_shape,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
splits,
|
||||
max_swizzle,
|
||||
raster_order_option,
|
||||
decomposition_mode,
|
||||
reduction_mode,
|
||||
reduction_warp_groups,
|
||||
barrier_bits,
|
||||
element_accumulator_bits,
|
||||
epilogue_subtile,
|
||||
num_accumulator_mtxs,
|
||||
cuda_adapter,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
|
||||
// Version of initialize_workspace that takes in as input the number of CTAs in the M and N dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
static cutlass::Status
|
||||
initialize_workspace(
|
||||
void* workspace,
|
||||
cudaStream_t stream,
|
||||
dim3 problem_blocks,
|
||||
uint32_t k_tiles_per_output_tile,
|
||||
GemmCoord tile_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int splits,
|
||||
int max_swizzle,
|
||||
RasterOrderOptions raster_order_option,
|
||||
DecompositionMode decomposition_mode,
|
||||
ReductionMode reduction_mode,
|
||||
uint32_t reduction_warp_groups,
|
||||
uint32_t barrier_bits,
|
||||
uint32_t element_accumulator_bits,
|
||||
uint32_t epilogue_subtile = 1,
|
||||
uint32_t num_accumulator_mtxs = 1,
|
||||
CudaHostAdapter *cuda_adapter = nullptr,
|
||||
uint32_t ktile_start_alignment_count = 1
|
||||
) {
|
||||
return UnderlyingStreamKParams::initialize_workspace(
|
||||
workspace,
|
||||
stream,
|
||||
problem_blocks,
|
||||
k_tiles_per_output_tile,
|
||||
tile_shape,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
splits,
|
||||
max_swizzle,
|
||||
raster_order_option,
|
||||
decomposition_mode,
|
||||
reduction_mode,
|
||||
reduction_warp_groups,
|
||||
barrier_bits,
|
||||
element_accumulator_bits,
|
||||
epilogue_subtile,
|
||||
num_accumulator_mtxs,
|
||||
cuda_adapter,
|
||||
ktile_start_alignment_count
|
||||
);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
// Parameters for SM100 persistent group scheduler (only used for Grouped Gemms)
|
||||
template<class ProblemShape>
|
||||
struct PersistentTileSchedulerSm100GroupParams {
|
||||
|
||||
using UnderlyingSm90Params = PersistentTileSchedulerSm90GroupParams<ProblemShape>;
|
||||
using RasterOrder = typename UnderlyingSm90Params::RasterOrder;
|
||||
using RasterOrderOptions = typename UnderlyingSm90Params::RasterOrderOptions;
|
||||
|
||||
UnderlyingSm90Params params_sm90_{};
|
||||
|
||||
// Version of initialize that takes in as input the number of CTAs in the M and N and L dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
void
|
||||
initialize(
|
||||
dim3 problem_blocks,
|
||||
int32_t groups,
|
||||
ProblemShape* problem_shapes,
|
||||
ProblemShape const* host_problem_shapes,
|
||||
GemmCoord cta_shape,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo const& hw_info,
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option
|
||||
) {
|
||||
|
||||
params_sm90_.initialize(
|
||||
problem_blocks,
|
||||
groups,
|
||||
problem_shapes,
|
||||
host_problem_shapes,
|
||||
cta_shape,
|
||||
cluster_shape,
|
||||
hw_info,
|
||||
max_swizzle_size,
|
||||
raster_order_option
|
||||
);
|
||||
}
|
||||
|
||||
// Version of get_tiled_cta_shape_mnl that takes in as input the number of CTAs in the M and N dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
CUTLASS_HOST_DEVICE
|
||||
static dim3
|
||||
get_tiled_cta_shape_mnl(GemmCoord cluster_shape, uint32_t cta_m, uint32_t cta_n) {
|
||||
return UnderlyingSm90Params::get_tiled_cta_shape_mnl(cluster_shape, cta_m, cta_n);
|
||||
}
|
||||
|
||||
// Version of get_grid_shape that takes in as input the number of CTAs in the M and N and L dimensions.
|
||||
// This is useful for calculating the tiled shape when a mode of problem and/or CTA shape has rank > 1,
|
||||
// for which using CuTe algebra for calculating tile shapes is easiest.
|
||||
CUTLASS_HOST_DEVICE static
|
||||
dim3
|
||||
get_grid_shape(
|
||||
dim3 problem_blocks,
|
||||
GemmCoord cluster_shape,
|
||||
KernelHardwareInfo hw_info,
|
||||
int max_swizzle_size,
|
||||
RasterOrderOptions raster_order_option,
|
||||
bool truncate_by_problem_size = true,
|
||||
bool is_static_cluster_shape = false) {
|
||||
|
||||
int const sm_count = hw_info.sm_count;
|
||||
|
||||
// Round up to nearest multiple of swizzle_size along each mode
|
||||
auto log_swizzle_size = get_log_swizzle_size(problem_blocks.x, problem_blocks.y, max_swizzle_size);
|
||||
auto problem_blocks_m = round_up(problem_blocks.x, (1 << log_swizzle_size) * cluster_shape.m());
|
||||
auto problem_blocks_n = round_up(problem_blocks.y, (1 << log_swizzle_size) * cluster_shape.n());
|
||||
|
||||
int problem_blocks_total = problem_blocks_m * problem_blocks_n * problem_blocks.z;
|
||||
|
||||
RasterOrder raster_order = get_rasterization_order(
|
||||
problem_blocks_m,
|
||||
problem_blocks_n,
|
||||
raster_order_option
|
||||
);
|
||||
|
||||
dim3 launch_grid;
|
||||
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid = dim3(cluster_shape.m(), 1, 1);
|
||||
}
|
||||
else {
|
||||
launch_grid = dim3(1, cluster_shape.n(), 1);
|
||||
}
|
||||
|
||||
auto possibly_truncate = [&](int x, int y) {
|
||||
if (truncate_by_problem_size) {
|
||||
return platform::min(x, y);
|
||||
}
|
||||
else {
|
||||
return x;
|
||||
}
|
||||
};
|
||||
|
||||
if (is_static_cluster_shape) {
|
||||
// The else path is generic, however, we can avoid some divs if we know cluster size is 1
|
||||
auto cluster_size = cluster_shape.m() * cluster_shape.n();
|
||||
if (cluster_size == 1) {
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid.y = possibly_truncate(sm_count, problem_blocks_total);
|
||||
}
|
||||
else {
|
||||
launch_grid.x = possibly_truncate(sm_count, problem_blocks_total);
|
||||
}
|
||||
}
|
||||
else {
|
||||
constexpr int max_sm_per_gpc = 20;
|
||||
int cta_per_device = get_max_cta_occupancy(max_sm_per_gpc, cluster_shape, sm_count);
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid.y = possibly_truncate(
|
||||
cta_per_device / cluster_shape.m(),
|
||||
problem_blocks_total / cluster_shape.m());
|
||||
}
|
||||
else {
|
||||
launch_grid.x = possibly_truncate(
|
||||
cta_per_device / cluster_shape.n(),
|
||||
problem_blocks_total / cluster_shape.n());
|
||||
}
|
||||
CUTLASS_TRACE_HOST("get_grid_shape(): Proposed GridDims by the scheduler using heuristics = "
|
||||
"(" << launch_grid.x << ", " << launch_grid.y << ", " << launch_grid.z << ")\n");
|
||||
}
|
||||
}
|
||||
else {
|
||||
// With preferred clusters, we can launch the largest possible persistent grid (rounded up to cluster dims)
|
||||
if (raster_order == RasterOrder::AlongN) {
|
||||
launch_grid.y = ((possibly_truncate(sm_count, problem_blocks_total) / cluster_shape.m()) / cluster_shape.n()) * cluster_shape.n();
|
||||
}
|
||||
else {
|
||||
launch_grid.x = ((possibly_truncate(sm_count, problem_blocks_total) / cluster_shape.n()) / cluster_shape.m()) * cluster_shape.m();
|
||||
}
|
||||
CUTLASS_TRACE_HOST("get_grid_shape(): Proposed GridDims by the scheduler using preferred clusters = "
|
||||
"(" << launch_grid.x << ", " << launch_grid.y << ", " << launch_grid.z << ")\n");
|
||||
}
|
||||
return launch_grid;
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static int32_t
|
||||
get_log_swizzle_size(int problem_ctas_m, int problem_ctas_n, int max_swizzle_size) {
|
||||
return UnderlyingSm90Params::get_log_swizzle_size(problem_ctas_m, problem_ctas_n, max_swizzle_size);
|
||||
}
|
||||
|
||||
CUTLASS_HOST_DEVICE
|
||||
static RasterOrder
|
||||
get_rasterization_order(
|
||||
uint32_t tiles_m,
|
||||
uint32_t tiles_n,
|
||||
RasterOrderOptions raster_order_option
|
||||
) {
|
||||
return UnderlyingSm90Params::get_rasterization_order(tiles_m, tiles_n, raster_order_option);
|
||||
}
|
||||
};
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
|
||||
} // namespace detail
|
||||
} // namespace kernel
|
||||
} // namespace gemm
|
||||
|
||||
Reference in New Issue
Block a user