CUTLASS 3.8 Release (#2059)

* CUTLASS 3.8 Release

* update

* Update README.md

* Revert "Update README.md"

This reverts commit b353e36fe83e0815f99b44e46c0c95494c44726b.

* update

* update

---------

Co-authored-by: Haicheng Wu <57973641+hwu36@users.noreply.github.com>
Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
This commit is contained in:
mihir-awatramani
2025-01-25 02:44:06 -05:00
committed by GitHub
co-authored by Haicheng Wu Haicheng Wu
parent 9eb01fa0b0
commit 389e493055
290 changed files with 91222 additions and 291 deletions
@@ -0,0 +1,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 ? &params.tma_load_a_fallback : &params.tma_load_a;
observed_tma_load_b_ = is_fallback_cluster ? &params.tma_load_b_fallback : &params.tma_load_b;
}
else {
observed_tma_load_a_ = &params.tma_load_a;
observed_tma_load_b_ = &params.tma_load_b;
}
}
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 ? &params.tma_load_a_fallback : &params.tma_load_a;
observed_tma_load_b_ = is_fallback_cluster ? &params.tma_load_b_fallback : &params.tma_load_b;
}
else {
observed_tma_load_a_ = &params.tma_load_a;
observed_tma_load_b_ = &params.tma_load_b;
}
}
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[] = {&params};
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 {
+170
View File
@@ -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
View File
@@ -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
View File
@@ -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