v3.8.0 update (#2082)
* 3.8 update * fix Markus' name --------- Co-authored-by: yuzhai <yuzhai@nvidia.com>
This commit is contained in:
@@ -0,0 +1,278 @@
|
||||
/***************************************************************************************************
|
||||
* Copyright (c) 2024 - 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
|
||||
* SPDX-License-Identifier: BSD-3-Clause
|
||||
*
|
||||
* Redistribution and use in source and binary forms, with or without
|
||||
* modification, are permitted provided that the following conditions are met:
|
||||
*
|
||||
* 1. Redistributions of source code must retain the above copyright notice, this
|
||||
* list of conditions and the following disclaimer.
|
||||
*
|
||||
* 2. Redistributions in binary form must reproduce the above copyright notice,
|
||||
* this list of conditions and the following disclaimer in the documentation
|
||||
* and/or other materials provided with the distribution.
|
||||
*
|
||||
* 3. Neither the name of the copyright holder nor the names of its
|
||||
* contributors may be used to endorse or promote products derived from
|
||||
* this software without specific prior written permission.
|
||||
*
|
||||
* THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
|
||||
* AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
|
||||
* IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
|
||||
* DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
|
||||
* FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
|
||||
* DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
|
||||
* SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
|
||||
* CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
|
||||
* OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
|
||||
* OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
|
||||
*
|
||||
**************************************************************************************************/
|
||||
#pragma once
|
||||
//
|
||||
|
||||
//
|
||||
|
||||
#include "cutlass/gemm/collective/builders/sm100_common.inl"
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace cutlass::gemm::collective {
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
namespace detail {
|
||||
|
||||
template<
|
||||
int CapacityBytes,
|
||||
class CtaTileShape_MNK,
|
||||
class TiledMma,
|
||||
class KernelScheduleType,
|
||||
UMMA::Major UmmaMajorA,
|
||||
int ComplexComponent = 1,
|
||||
int NumComputeMtxs = 3,
|
||||
int carveout_bytes
|
||||
>
|
||||
constexpr cute::tuple<int, int, int>
|
||||
sm100_compute_stage_count_or_override_fast_fp32(StageCountAutoCarveout<carveout_bytes> stage_count) {
|
||||
constexpr int CtaM = get<0>(CtaTileShape_MNK{});
|
||||
constexpr int CtaN = get<1>(CtaTileShape_MNK{});
|
||||
static_assert(CtaN <= 128, "Can't support CtaN>128 tiles");
|
||||
constexpr int CtaK = get<2>(CtaTileShape_MNK{});
|
||||
using AtomThrID = typename TiledMma::AtomThrID;
|
||||
// Detect 2x2 TMEM layout
|
||||
constexpr int TmemAccWordsPerDP = (CtaM == 64 && size(AtomThrID{}) == 2) ? CtaN/2 : CtaN;
|
||||
constexpr int TmemAWordsPerDP = ComplexComponent * NumComputeMtxs * CtaK / 2;
|
||||
constexpr bool IsAComputeinTmem = UmmaMajorA == cute::UMMA::Major::K && !cute::is_base_of_v<KernelTmaWarpSpecializedFastFP32SmemSm100, KernelScheduleType>;
|
||||
constexpr bool IsAComputeinSmem = !IsAComputeinTmem;
|
||||
constexpr int AccumulatorStageCount = (IsAComputeinTmem) ? (((TmemAccWordsPerDP * ComplexComponent == 128) ? 2 : 3) * ComplexComponent) : (512 / TmemAccWordsPerDP);
|
||||
|
||||
constexpr int SmemCapacityAfterMma2AccumCarveout = CapacityBytes - (carveout_bytes + AccumulatorStageCount * 32);
|
||||
|
||||
constexpr int TmemInAStageCount_Potential = (IsAComputeinTmem) ? (512 - AccumulatorStageCount * TmemAccWordsPerDP) / TmemAWordsPerDP : 10000;
|
||||
|
||||
constexpr auto load2transform_pipeline_bytes = sizeof(typename cutlass::PipelineTmaTransformAsync<1>::SharedStorage);
|
||||
constexpr auto a_bits = cute::sizeof_bits_v<float> * ComplexComponent;
|
||||
constexpr auto b_bits = cute::sizeof_bits_v<float> * ComplexComponent;
|
||||
constexpr int ab_stage_bytes =
|
||||
cutlass::bits_to_bytes(a_bits * size<0>(CtaTileShape_MNK{}) * size<2>(CtaTileShape_MNK{})) +
|
||||
cutlass::bits_to_bytes(b_bits * size<1>(CtaTileShape_MNK{}) / size(AtomThrID{}) * size<2>(CtaTileShape_MNK{})) +
|
||||
static_cast<int>(load2transform_pipeline_bytes);
|
||||
|
||||
constexpr auto transform2mma_pipeline_bytes = sizeof(typename cutlass::PipelineUmmaConsumerAsync<1>::SharedStorage);
|
||||
constexpr auto a_compute_bits = cute::sizeof_bits_v<cutlass::bfloat16_t> * ComplexComponent;
|
||||
constexpr auto b_compute_bits = cute::sizeof_bits_v<cutlass::bfloat16_t> * ComplexComponent * ComplexComponent;
|
||||
constexpr int ab_compute_stage_bytes =
|
||||
cutlass::bits_to_bytes(NumComputeMtxs * a_compute_bits * int(IsAComputeinSmem) * size<0>(CtaTileShape_MNK{}) * size<2>(CtaTileShape_MNK{})) + // If ACompute is in TMEM, Acompute buffer has 0 bytes.
|
||||
cutlass::bits_to_bytes(NumComputeMtxs * b_compute_bits * size<1>(CtaTileShape_MNK{}) / size(AtomThrID{}) * size<2>(CtaTileShape_MNK{})) +
|
||||
static_cast<int>(transform2mma_pipeline_bytes);
|
||||
|
||||
constexpr int ABComputeStageCount_Potential = SmemCapacityAfterMma2AccumCarveout / (ab_stage_bytes + ab_compute_stage_bytes);
|
||||
// The number of SMEM buffers for A, B. ACompute (if in SMEM), BCompute should be at least Transform2MmaStageCount
|
||||
constexpr int Transform2MmaStageCount = std::min(TmemInAStageCount_Potential, ABComputeStageCount_Potential);
|
||||
|
||||
constexpr int SmemCapacityAfterABComputeCarveout = SmemCapacityAfterMma2AccumCarveout - (Transform2MmaStageCount * ab_compute_stage_bytes);
|
||||
// Can we boost the number of buffers for A and B?
|
||||
constexpr int Load2TransformStageCount = SmemCapacityAfterABComputeCarveout / ab_stage_bytes;
|
||||
|
||||
static_assert(Load2TransformStageCount >= 2 && Transform2MmaStageCount >= 2 && AccumulatorStageCount >= 2, "Not enough SMEM or TMEM capacity for selected tile size");
|
||||
return cute::make_tuple(Load2TransformStageCount, Transform2MmaStageCount, AccumulatorStageCount);
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
|
||||
// FastFP (9xBF16) MMA kernels builder
|
||||
template <
|
||||
class GmemLayoutATag,
|
||||
int AlignmentA,
|
||||
class GmemLayoutBTag,
|
||||
int AlignmentB,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK, // The Cluster-level TileShape
|
||||
class ClusterShape_MNK,
|
||||
class StageCountType,
|
||||
class KernelScheduleType
|
||||
>
|
||||
struct CollectiveBuilder<
|
||||
arch::Sm100,
|
||||
arch::OpClassTensorOp,
|
||||
float, // ElementA
|
||||
GmemLayoutATag, // LayoutA
|
||||
AlignmentA,
|
||||
float, // ElementB
|
||||
GmemLayoutBTag, // LayoutB
|
||||
AlignmentB,
|
||||
ElementAccumulator,
|
||||
TileShape_MNK, // (MmaAtomShapeM, MmaAtomShapeN, TileK)
|
||||
ClusterShape_MNK, // Static cluster shape or dynamic (int, int, int)
|
||||
StageCountType,
|
||||
KernelScheduleType,
|
||||
cute::enable_if_t<
|
||||
(not cute::is_tuple<GmemLayoutATag>::value && not cute::is_tuple<GmemLayoutBTag>::value) &&
|
||||
(cute::is_base_of_v<KernelScheduleSm100FastFP32Gemm, KernelScheduleType>) &&
|
||||
((sizeof(float) * AlignmentA) % detail::tma_alignment_bytes == 0) &&
|
||||
((sizeof(float) * AlignmentB) % detail::tma_alignment_bytes == 0)>>
|
||||
{
|
||||
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>();
|
||||
|
||||
using ElementA = float;
|
||||
using ElementB = float;
|
||||
using ElementAMma = cutlass::bfloat16_t;
|
||||
using ElementBMma = cutlass::bfloat16_t;
|
||||
static constexpr int ScalingFactor = 8;
|
||||
|
||||
using TiledMma = decltype(detail::sm100_make_trivial_fastFP32_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator, TileShape_MNK, ClusterShape_MNK, UmmaMajorA, UmmaMajorB, ScalingFactor, KernelScheduleType>());
|
||||
using AtomThrID = typename TiledMma::AtomThrID;
|
||||
using AtomThrShapeMNK = Shape<decltype(shape<0>(typename TiledMma::ThrLayoutVMNK{})), _1, _1>;
|
||||
using CtaTileShape_MNK = decltype(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 SmemLayoutAtomA = decltype(cutlass::gemm::collective::detail::sm100_smem_selector<UmmaMajorA, ElementA,
|
||||
BlockTileA_M, BlockTileA_K>());
|
||||
// Take 3 compute buffers into account for swizzle selection
|
||||
using SmemLayoutAtomACompute = decltype(cutlass::gemm::collective::detail::sm100_smem_selector<UmmaMajorA, ElementAMma,
|
||||
BlockTileA_M, BlockTileA_K>());
|
||||
|
||||
// Input transform kernel can not use TMA 2SM instructions.
|
||||
using GmemTiledCopyA = decltype(detail::sm90_cluster_shape_to_tma_atom(cute::size<1>(ClusterShape_MNK{})));
|
||||
using SmemLayoutAtomPairA = cutlass::gemm::collective::detail::CollectiveMmaEmulatedLayoutAtomType<
|
||||
SmemLayoutAtomA, SmemLayoutAtomACompute>;
|
||||
|
||||
static constexpr int MMA_M = cute::size<0,0>(MmaShapeA_MK{});
|
||||
using CopyAtomPairA = cutlass::gemm::collective::detail::CollectiveMmaEmulatedCopyType<
|
||||
Copy_Atom<AutoVectorizingCopyWithAssumedAlignment<128>, ElementA>,
|
||||
cute::conditional_t<(UmmaMajorA == cute::UMMA::Major::K && !cute::is_base_of_v<KernelTmaWarpSpecializedFastFP32SmemSm100, KernelScheduleType>),
|
||||
cute::conditional_t<(MMA_M == 64 && size(AtomThrID{}) == 1), SM100_TMEM_STORE_16dp256b1x, SM100_TMEM_STORE_32dp32b8x>, // TS Implementation
|
||||
Copy_Atom<AutoVectorizingCopyWithAssumedAlignment<128>, ElementA>> // SS Implementation
|
||||
>;
|
||||
|
||||
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{}));
|
||||
|
||||
// Input transform kernel can not use TMA 2SM instructions.
|
||||
using GmemTiledCopyB = decltype(detail::sm90_cluster_shape_to_tma_atom(cute::size<0>(ClusterShape_MNK{})));
|
||||
|
||||
using SmemLayoutAtomB = decltype(cutlass::gemm::collective::detail::sm100_smem_selector<UmmaMajorB, ElementB,
|
||||
BlockTileB_N, BlockTileB_K>());
|
||||
// Take 3 compute buffers into account for swizzle selection
|
||||
using SmemLayoutAtomBCompute = decltype(cutlass::gemm::collective::detail::sm100_smem_selector<UmmaMajorB, ElementBMma,
|
||||
BlockTileB_N, BlockTileB_K>());
|
||||
|
||||
using SmemLayoutAtomPairB = cutlass::gemm::collective::detail::CollectiveMmaEmulatedLayoutAtomType<
|
||||
SmemLayoutAtomB, SmemLayoutAtomBCompute>;
|
||||
using CopyAtomPairB = cutlass::gemm::collective::detail::CollectiveMmaEmulatedCopyType<
|
||||
Copy_Atom<AutoVectorizingCopyWithAssumedAlignment<128>, ElementB>,
|
||||
Copy_Atom<AutoVectorizingCopyWithAssumedAlignment<128>, ElementBMma>
|
||||
>;
|
||||
|
||||
// SmemCarveout
|
||||
static constexpr int NumBandsToCompute = 5;
|
||||
static constexpr int AccPromotionInterval = 1;
|
||||
static constexpr int SchedulerPipelineStageCount = 3;
|
||||
static constexpr bool IsArrayOfPointersGemm = (cute::is_base_of_v<KernelScheduleSm100PtrArrayFastFP32Gemm, KernelScheduleType>);
|
||||
|
||||
// CLCPipeline = PipelineCLCFetchAsync
|
||||
static constexpr auto CLCPipelineStorage = sizeof(typename cutlass::PipelineCLCFetchAsync<SchedulerPipelineStageCount, ClusterShape_MNK>::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 = sizeof(uint32_t);
|
||||
// Tensormap Storage
|
||||
static constexpr size_t TensorMapStorage = IsArrayOfPointersGemm ? sizeof(cute::TmaDescriptor) * 2 /* for A and B */ : 0;
|
||||
|
||||
// Smem usage that's not part of CollectiveEpilogue::SharedStorage & CollectiveMainloop::SharedStorage
|
||||
static constexpr auto KernelSmemCarveout = static_cast<int>( CLCPipelineStorage +
|
||||
CLCResponseStorage +
|
||||
CLCThrottlePipelineStorage +
|
||||
TmemDeallocStorage +
|
||||
TmemBasePtrsStorage +
|
||||
TensorMapStorage);
|
||||
|
||||
// Reduce SMEM capacity available for buffers considering extra B smem and barrier smem allocations
|
||||
static constexpr int Sm100ReducedSmemCapacityBytes = detail::sm100_smem_capacity_bytes - KernelSmemCarveout;
|
||||
static constexpr auto stage_info = cutlass::gemm::collective::detail::sm100_compute_stage_count_or_override_fast_fp32<
|
||||
Sm100ReducedSmemCapacityBytes, CtaTileShape_MNK, TiledMma, KernelScheduleType, UmmaMajorA>(StageCountType{});
|
||||
|
||||
static constexpr int Load2TransformPipelineStageCount = get<0>(stage_info);
|
||||
static constexpr int Transform2MmaPipelineStageCount = get<1>(stage_info);
|
||||
static constexpr int AccumulatorPipelineStageCount = get<2>(stage_info);
|
||||
|
||||
using AccumulatorCopyAtom = cute::SM100_TMEM_LOAD_32dp32b32x;
|
||||
|
||||
using DispatchPolicy = cute::conditional_t<IsArrayOfPointersGemm,
|
||||
cutlass::gemm::MainloopSm100ArrayTmaUmmaWarpSpecializedFastF32<
|
||||
Load2TransformPipelineStageCount,
|
||||
Transform2MmaPipelineStageCount,
|
||||
SchedulerPipelineStageCount,
|
||||
AccumulatorPipelineStageCount,
|
||||
NumBandsToCompute,
|
||||
ScalingFactor,
|
||||
AccPromotionInterval,
|
||||
ClusterShape_MNK,
|
||||
AccumulatorCopyAtom>,
|
||||
cutlass::gemm::MainloopSm100TmaUmmaWarpSpecializedFastF32<
|
||||
Load2TransformPipelineStageCount,
|
||||
Transform2MmaPipelineStageCount,
|
||||
SchedulerPipelineStageCount,
|
||||
AccumulatorPipelineStageCount,
|
||||
NumBandsToCompute,
|
||||
ScalingFactor,
|
||||
AccPromotionInterval,
|
||||
ClusterShape_MNK,
|
||||
AccumulatorCopyAtom>
|
||||
>;
|
||||
using CollectiveOp = cutlass::gemm::collective::CollectiveMma<
|
||||
DispatchPolicy,
|
||||
TileShape_MNK,
|
||||
ElementA,
|
||||
cutlass::gemm::TagToStrideA_t<GmemLayoutATag>,
|
||||
ElementB,
|
||||
cutlass::gemm::TagToStrideB_t<GmemLayoutBTag>,
|
||||
TiledMma,
|
||||
GmemTiledCopyA,
|
||||
SmemLayoutAtomPairA,
|
||||
CopyAtomPairA,
|
||||
cute::identity,
|
||||
GmemTiledCopyB,
|
||||
SmemLayoutAtomPairB,
|
||||
CopyAtomPairB,
|
||||
cute::identity
|
||||
>;
|
||||
};
|
||||
|
||||
} // namespace cutlass::gemm::collective
|
||||
@@ -71,7 +71,7 @@ template <
|
||||
>
|
||||
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
|
||||
// For MXF8F6F4 MMA, 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)
|
||||
@@ -386,7 +386,7 @@ select_instr() {
|
||||
}
|
||||
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
|
||||
// Fp4 can be mixed with FP6, Fp8 with MMA.MXF8F6F4 only
|
||||
return detail::blockscaled::BlockScaledInstr::MXF4F6F8;
|
||||
}
|
||||
else if constexpr (sizeof_bits_v<ElementA> == 4 && sizeof_bits_v<ElementB> == 4) {
|
||||
@@ -400,7 +400,7 @@ select_instr() {
|
||||
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");
|
||||
"Only MXF4 support with non-TN and MMA.MXF8F6F4.");
|
||||
return detail::blockscaled::BlockScaledInstr::MXF4F6F8;
|
||||
}
|
||||
}
|
||||
@@ -636,7 +636,7 @@ struct CollectiveBuilder<
|
||||
|
||||
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");
|
||||
static_assert(UseMxf8f6f4 || (cutlass::gemm::detail::is_k_major_A<GmemLayoutATag>() && cutlass::gemm::detail::is_k_major_B<GmemLayoutBTag>()), "Only MMA.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>());
|
||||
|
||||
@@ -477,6 +477,94 @@ sm100_make_trivial_tiled_mma() {
|
||||
}
|
||||
}
|
||||
|
||||
template<
|
||||
class ElementAMma,
|
||||
class ElementBMma,
|
||||
class ElementAccumulator,
|
||||
class TileShape_MNK,
|
||||
class ClusterShape_MNK,
|
||||
UMMA::Major UmmaMajorA,
|
||||
UMMA::Major UmmaMajorB,
|
||||
int Scale,
|
||||
class KernelScheduleType
|
||||
>
|
||||
constexpr auto
|
||||
sm100_make_trivial_fastFP32_tiled_mma() {
|
||||
// MMA_2SM requested
|
||||
if constexpr (cute::is_base_of_v<KernelSchedule2Sm, KernelScheduleType> ) {
|
||||
using AtomLayout_MNK = decltype(make_layout(shape_div(ClusterShape_MNK{}, Shape<_2,_1,_1>{})));
|
||||
constexpr int M = cute::size<0>(TileShape_MNK{});
|
||||
constexpr int N = cute::size<1>(TileShape_MNK{});
|
||||
if constexpr (UmmaMajorA == cute::UMMA::Major::K && !cute::is_base_of_v<KernelTmaWarpSpecializedFastFP32SmemSm100, KernelScheduleType>) {
|
||||
return make_tiled_mma(cute::SM100_MMA_F16BF16_2x1SM_TS_SCALED<ElementAMma, ElementBMma, ElementAccumulator,
|
||||
M, N, UmmaMajorA, UmmaMajorB, Scale>{});
|
||||
}
|
||||
else { // If A needs to be transposed by MMA, fall back to SMEM from A MMA instructions
|
||||
return make_tiled_mma(cute::SM100_MMA_F16BF16_2x1SM_SS_SCALED<ElementAMma, ElementBMma, ElementAccumulator,
|
||||
M, N, UmmaMajorA, UmmaMajorB, Scale>{});
|
||||
}
|
||||
}
|
||||
// MMA_1SM requested
|
||||
else if constexpr (cute::is_base_of_v<KernelSchedule1Sm, KernelScheduleType> ) {
|
||||
// using AtomLayout_MNK = Layout<ClusterShape_MNK>;
|
||||
constexpr int M = cute::size<0>(TileShape_MNK{});
|
||||
constexpr int N = cute::size<1>(TileShape_MNK{});
|
||||
if constexpr (UmmaMajorA == cute::UMMA::Major::K && !cute::is_base_of_v<KernelTmaWarpSpecializedFastFP32SmemSm100, KernelScheduleType>) {
|
||||
return make_tiled_mma(cute::SM100_MMA_F16BF16_TS_SCALED<ElementAMma, ElementBMma, ElementAccumulator,
|
||||
M, N, UmmaMajorA, UmmaMajorB, Scale>{});
|
||||
}
|
||||
else { // If A needs to be transposed by MMA, fall back to SMEM from A MMA instructions
|
||||
return make_tiled_mma(cute::SM100_MMA_F16BF16_SS_SCALED<ElementAMma, ElementBMma, ElementAccumulator,
|
||||
M, N, UmmaMajorA, UmmaMajorB, Scale>{});
|
||||
}
|
||||
}
|
||||
else if constexpr (cute::is_same_v<KernelScheduleType, KernelScheduleSm100FastFP32Gemm> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedFastFP32SmemSm100> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelScheduleSm100PtrArrayFastFP32Gemm> ||
|
||||
cute::is_same_v<KernelScheduleType, KernelTmaWarpSpecializedPtrArrayFastFP32SmemSm100>) {
|
||||
// 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::get<0>(ClusterShape_MNK{}) % 2 == 0 &&
|
||||
(cute::get<0>(TileShape_MNK{}) / cute::get<0>(ClusterShape_MNK{})) % 64 == 0) {
|
||||
if constexpr (!cute::is_base_of_v<KernelTmaWarpSpecializedFastFP32SmemSm100, KernelScheduleType>) {
|
||||
return sm100_make_trivial_fastFP32_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator, TileShape_MNK,
|
||||
ClusterShape_MNK, UmmaMajorA, UmmaMajorB, Scale, KernelTmaWarpSpecialized2SmFastFP32Sm100>();
|
||||
}
|
||||
else {
|
||||
return sm100_make_trivial_fastFP32_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator, TileShape_MNK,
|
||||
ClusterShape_MNK, UmmaMajorA, UmmaMajorB, Scale, KernelTmaWarpSpecialized2SmFastFP32SmemSm100>();
|
||||
}
|
||||
}
|
||||
else {
|
||||
if constexpr (!cute::is_base_of_v<KernelTmaWarpSpecializedFastFP32SmemSm100, KernelScheduleType>) {
|
||||
return sm100_make_trivial_fastFP32_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator, TileShape_MNK,
|
||||
ClusterShape_MNK, UmmaMajorA, UmmaMajorB, Scale, KernelTmaWarpSpecialized1SmFastFP32Sm100>();
|
||||
}
|
||||
else {
|
||||
return sm100_make_trivial_fastFP32_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator, TileShape_MNK,
|
||||
ClusterShape_MNK, UmmaMajorA, UmmaMajorB, Scale, KernelTmaWarpSpecialized1SmFastFP32SmemSm100>();
|
||||
}
|
||||
}
|
||||
}
|
||||
// Dynamic cluster shape means we cannot assume we can use 2SM MMA
|
||||
else {
|
||||
if constexpr (!cute::is_base_of_v<KernelTmaWarpSpecializedFastFP32SmemSm100, KernelScheduleType>) {
|
||||
return sm100_make_trivial_fastFP32_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator, TileShape_MNK,
|
||||
ClusterShape_MNK, UmmaMajorA, UmmaMajorB, Scale, KernelTmaWarpSpecialized1SmFastFP32Sm100>();
|
||||
}
|
||||
else {
|
||||
return sm100_make_trivial_fastFP32_tiled_mma<ElementAMma, ElementBMma, ElementAccumulator, TileShape_MNK,
|
||||
ClusterShape_MNK, UmmaMajorA, UmmaMajorB, Scale, KernelTmaWarpSpecialized1SmFastFP32SmemSm100>();
|
||||
}
|
||||
}
|
||||
}
|
||||
else {
|
||||
static_assert(cutlass::detail::dependent_false<TileShape_MNK> == 0,
|
||||
"Unsupported policy for SM100 collective builder.");
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* @brief Check for U4_UNPACK_U8, U6_UNPACK_U8 alignment requirement
|
||||
@@ -547,22 +635,22 @@ template <class ElementA, int AlignmentA, class ElementB, int AlignmentB, class
|
||||
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;
|
||||
constexpr bool is_f8f6f4_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);
|
||||
return ((cute::sizeof_bits_v<ElementA> * AlignmentA) % cutlass::detail::get_input_alignment_bits<ElementA, is_f8f6f4_subbytes>() == 0) &&
|
||||
((cute::sizeof_bits_v<ElementB> * AlignmentB) % cutlass::detail::get_input_alignment_bits<ElementB, is_f8f6f4_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) &&
|
||||
constexpr bool is_mxf8f6f4_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);
|
||||
return ((cute::sizeof_bits_v<ElementA> * AlignmentA) % cutlass::detail::get_input_alignment_bits<ElementA, is_mxf8f6f4_subbytes>() == 0) &&
|
||||
((cute::sizeof_bits_v<ElementB> * AlignmentB) % cutlass::detail::get_input_alignment_bits<ElementB, is_mxf8f6f4_subbytes>() == 0);
|
||||
}
|
||||
|
||||
} // namespace detail
|
||||
|
||||
@@ -82,7 +82,7 @@ template<
|
||||
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 F8/F6/F4 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)
|
||||
@@ -253,7 +253,9 @@ struct CollectiveBuilder<
|
||||
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 constexpr uint32_t AccumulatorPipelineStageCount = (is_2sm || (!is_2sm && size(shape<0,0>(MmaShapeA_MK{}) > 64))) ?
|
||||
TotalTmem / (cute::size<0>(CtaTileShape_MNK{}) * cute::size<1>(CtaTileShape_MNK{}))
|
||||
: (Sm100TmemCapacityColumns / cute::size<1>(CtaTileShape_MNK{})) * 2; // 1SM MMA_M = 64 case
|
||||
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.
|
||||
|
||||
@@ -261,8 +261,9 @@ struct CollectiveBuilder<
|
||||
using SmemLayoutAtomB = decltype(detail::ss_smem_selector<
|
||||
GmmaMajorB, ElementBMma, decltype(cute::get<1>(TileShape_MNK{})), decltype(cute::get<2>(TileShape_MNK{}))>());
|
||||
|
||||
static constexpr int Sm90ReducedSmemCapacityBytes =
|
||||
detail::sm90_smem_capacity_bytes;
|
||||
static constexpr size_t TensorMapStorage = IsArrayOfPointersGemm ? sizeof(cute::TmaDescriptor) * 2 /* for A and B */ : 0;
|
||||
static constexpr int KernelSmemCarveout = static_cast<int>(TensorMapStorage);
|
||||
static constexpr int Sm90ReducedSmemCapacityBytes = detail::sm90_smem_capacity_bytes - KernelSmemCarveout;
|
||||
|
||||
static constexpr int PipelineStages = detail::compute_stage_count_or_override<Sm90ReducedSmemCapacityBytes,
|
||||
ElementAMma, ElementBMma, TileShape_MNK>(StageCountType{});
|
||||
@@ -368,7 +369,12 @@ public:
|
||||
return t;
|
||||
}
|
||||
else {
|
||||
if constexpr (cute::is_pointer_v<T>) {
|
||||
return &cute::stride(*t);
|
||||
}
|
||||
else {
|
||||
return cute::stride(t);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -441,14 +447,20 @@ public:
|
||||
static constexpr int Sm90ReducedSmemCapacityBytes = detail::sm90_smem_capacity_bytes - KernelSmemCarveout;
|
||||
|
||||
static constexpr int PipelineStages = IsMixedInput ?
|
||||
( IsArrayOfPointersGemm ?
|
||||
detail::compute_stage_count_or_override_single_affine_transformed_input<Sm90ReducedSmemCapacityBytes,
|
||||
RealElementA, RealElementB, ElementScale, ElementZero, TileShape_MNK, StageCountType::bytes, SmemAlignment>(StageCountType{}) :
|
||||
detail::compute_stage_count_or_override_single_affine_transformed_input<detail::sm90_smem_capacity_bytes,
|
||||
RealElementA, RealElementB, ElementScale, ElementZero, TileShape_MNK, StageCountType::bytes, SmemAlignment>(StageCountType{})
|
||||
)
|
||||
: detail::compute_stage_count_or_override<detail::sm90_smem_capacity_bytes,
|
||||
ElementAMma, ElementBMma, TileShape_MNK, StageCountType::bytes, SmemAlignment>(StageCountType{});
|
||||
|
||||
using DispatchPolicy = cute::conditional_t<IsMixedInput,
|
||||
MainloopSm90TmaGmmaRmemAWarpSpecializedMixedInput<PipelineStages, ClusterShape_MNK, KernelScheduleType>
|
||||
, MainloopSm90TmaGmmaRmemAWarpSpecialized<PipelineStages, ClusterShape_MNK, KernelScheduleType>>;
|
||||
cute::conditional_t<IsArrayOfPointersGemm,
|
||||
MainloopSm90ArrayTmaGmmaWarpSpecializedMixedInput<PipelineStages, ClusterShape_MNK, KernelScheduleType>,
|
||||
MainloopSm90TmaGmmaRmemAWarpSpecializedMixedInput<PipelineStages, ClusterShape_MNK, KernelScheduleType>>,
|
||||
MainloopSm90TmaGmmaRmemAWarpSpecialized<PipelineStages, ClusterShape_MNK, KernelScheduleType>>;
|
||||
|
||||
using SmemCopyAtomA = cute::conditional_t<SwapAB, void, Copy_Atom<cute::AutoVectorizingCopy, ElementA>>;
|
||||
using SmemCopyAtomB = cute::conditional_t<SwapAB, Copy_Atom<cute::AutoVectorizingCopy, ElementB>, void>;
|
||||
|
||||
@@ -71,15 +71,15 @@ struct Sm90GemmSparseConfig {
|
||||
using ElementEMmaSparsity = Int<ElementEMma::sparsity>;
|
||||
|
||||
// MMA type
|
||||
static constexpr bool IsQmma = cute::is_same_v<ElementAMmaRaw, float_e4m3_t> && ElementAMmaSparsity{} == _2{} ||
|
||||
static constexpr bool IsF8 = cute::is_same_v<ElementAMmaRaw, float_e4m3_t> && ElementAMmaSparsity{} == _2{} ||
|
||||
cute::is_same_v<ElementAMmaRaw, float_e5m2_t> && ElementAMmaSparsity{} == _2{};
|
||||
static constexpr bool IsImma = cute::is_same_v<ElementAMmaRaw, int8_t> && ElementAMmaSparsity{} == _2{} ||
|
||||
static constexpr bool IsI8 = cute::is_same_v<ElementAMmaRaw, int8_t> && ElementAMmaSparsity{} == _2{} ||
|
||||
cute::is_same_v<ElementAMmaRaw, uint8_t> && ElementAMmaSparsity{} == _2{};
|
||||
static constexpr bool IsHmma = cute::is_same_v<ElementAMmaRaw, half_t> && ElementAMmaSparsity{} == _2{} ||
|
||||
static constexpr bool IsF16BF16 = cute::is_same_v<ElementAMmaRaw, half_t> && ElementAMmaSparsity{} == _2{} ||
|
||||
cute::is_same_v<ElementAMmaRaw, bfloat16_t> && ElementAMmaSparsity{} == _2{};
|
||||
static constexpr bool IsTfmma = cute::is_same_v<ElementAMmaRaw, tfloat32_t> && ElementAMmaSparsity{} == _2{} ||
|
||||
static constexpr bool IsTF32 = cute::is_same_v<ElementAMmaRaw, tfloat32_t> && ElementAMmaSparsity{} == _2{} ||
|
||||
cute::is_same_v<ElementAMmaRaw, float> && ElementAMmaSparsity{} == _2{};
|
||||
static_assert(int(IsQmma) + int(IsImma) + int(IsHmma) + int(IsTfmma) == 1, "Ambigious Input Type Config (failed to choose MMA type)");
|
||||
static_assert(int(IsF8) + int(IsI8) + int(IsF16BF16) + int(IsTF32) == 1, "Ambigious Input Type Config (failed to choose MMA type)");
|
||||
|
||||
// Number of ElementARaw stored in ElementAMmaRaw. For Hopper this is always 1.
|
||||
using ElemsARawPerElementAMmaRaw = _1;
|
||||
@@ -89,12 +89,12 @@ struct Sm90GemmSparseConfig {
|
||||
static_assert(ElementASparsity{} == _2{}, "ElementASparsity must be 2 for Hopper Sparse Gemm");
|
||||
|
||||
// Logical/Physical ElementA per Chunk
|
||||
using LogicalElemsAPerChunk = conditional_t<IsTfmma, _2, _4>;
|
||||
using LogicalElemsAPerChunk = conditional_t<IsTF32, _2, _4>;
|
||||
using PhysicalElemsAPerChunk = Int<LogicalElemsAPerChunk{} / ElementASparsity{}>;
|
||||
|
||||
// Metadata Bits
|
||||
using ElementEBitsPerChunk = _4;
|
||||
using ElementEBitsPerElementAMma = cute::conditional_t<IsTfmma, _4, _2>;
|
||||
using ElementEBitsPerElementAMma = cute::conditional_t<IsTF32, _4, _2>;
|
||||
|
||||
// Metadata Layout. Unit in corresbonding logical elements.
|
||||
// Basic metadata block is (16,64) for 8-bit, (16,32) for 16-bit, (16,16) for 32-bit data types.
|
||||
@@ -114,8 +114,8 @@ struct Sm90GemmSparseConfig {
|
||||
using TensorEAtom_8bit = decltype(make_ordered_layout(Shape<_64,MinTileShapeK>{},
|
||||
Step < _1, _0>{}));
|
||||
|
||||
using TensorEAtom = cute::conditional_t<(IsQmma || IsImma), TensorEAtom_8bit,
|
||||
cute::conditional_t<IsTfmma, TensorEAtom_32bit,
|
||||
using TensorEAtom = cute::conditional_t<(IsF8 || IsI8), TensorEAtom_8bit,
|
||||
cute::conditional_t<IsTF32, TensorEAtom_32bit,
|
||||
TensorEAtom_16bit>>;
|
||||
|
||||
// Logical elems that construct the atomK for tensorE/A.
|
||||
|
||||
@@ -40,8 +40,9 @@
|
||||
#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"
|
||||
#include "cutlass/gemm/collective/builders/sm100_umma_builder.inl"
|
||||
#include "cutlass/gemm/collective/builders/sm100_9xBF16_umma_builder.inl"
|
||||
#include "cutlass/gemm/collective/builders/sm100_blockscaled_umma_builder.inl"
|
||||
#endif
|
||||
|
||||
/////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
@@ -46,11 +46,16 @@
|
||||
#include "cutlass/gemm/collective/sm90_sparse_mma_tma_gmma_ss_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_sparse_mma_tma_gmma_ss_warpspecialized_fp8.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_array_tma_gmma_ss_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm90_mma_array_tma_gmma_rs_warpspecialized_mixed_input.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"
|
||||
|
||||
#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_mma_warpspecialized_emulated.hpp"
|
||||
#include "cutlass/gemm/collective/sm100_mma_array_warpspecialized_emulated.hpp"
|
||||
|
||||
#include "cutlass/gemm/collective/sm100_blockscaled_mma_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/collective/sm100_blockscaled_mma_array_warpspecialized.hpp"
|
||||
#endif // !defined(__CUDACC_RTC__)
|
||||
|
||||
@@ -682,11 +682,11 @@ struct CollectiveMma<
|
||||
auto mSFB_nkl = [=](){
|
||||
if constexpr (IsCtaN192) {
|
||||
Tensor mSFB_tmp = observed_tma_load_sfb_->get_tma_tensor(shape(layout_SFB));
|
||||
auto x = stride<0,2>(mSFB_tmp);
|
||||
auto y = ceil_div(shape<0,2>(mSFB_tmp), 4);
|
||||
auto new_shape = make_shape (make_shape( shape<0,0>(mSFB_tmp), shape<0,1>(mSFB_tmp),
|
||||
auto x = stride<0,1>(mSFB_tmp);
|
||||
auto y = ceil_div(shape<0,1>(mSFB_tmp), 4);
|
||||
auto new_shape = make_shape (make_shape( shape<0,0>(mSFB_tmp),
|
||||
make_shape( make_shape(_2{}, _2{}), y)), shape<1>(mSFB_tmp), shape<2>(mSFB_tmp));
|
||||
auto new_stride = make_stride(make_stride(stride<0,0>(mSFB_tmp), stride<0,1>(mSFB_tmp),
|
||||
auto new_stride = make_stride(make_stride(stride<0,0>(mSFB_tmp),
|
||||
make_stride(make_stride( x, x), x*3)), stride<1>(mSFB_tmp), stride<2>(mSFB_tmp));
|
||||
return make_tensor(mSFB_tmp.data(), make_layout(new_shape, new_stride));
|
||||
}
|
||||
|
||||
@@ -717,11 +717,11 @@ struct CollectiveMma<
|
||||
auto mSFB_nkl = [=](){
|
||||
if constexpr (IsCtaN192) {
|
||||
Tensor mSFB_tmp = observed_tma_load_sfb_->get_tma_tensor(shape(layout_SFB_));
|
||||
auto x = stride<0,2>(mSFB_tmp);
|
||||
auto y = ceil_div(shape<0,2>(mSFB_tmp), 4);
|
||||
auto new_shape = make_shape (make_shape( shape<0,0>(mSFB_tmp), shape<0,1>(mSFB_tmp),
|
||||
auto x = stride<0,1>(mSFB_tmp);
|
||||
auto y = ceil_div(shape<0,1>(mSFB_tmp), 4);
|
||||
auto new_shape = make_shape (make_shape( shape<0,0>(mSFB_tmp),
|
||||
make_shape( make_shape(_2{}, _2{}), y)), shape<1>(mSFB_tmp), shape<2>(mSFB_tmp));
|
||||
auto new_stride = make_stride(make_stride(stride<0,0>(mSFB_tmp), stride<0,1>(mSFB_tmp),
|
||||
auto new_stride = make_stride(make_stride(stride<0,0>(mSFB_tmp),
|
||||
make_stride(make_stride( x, x), x*3)), stride<1>(mSFB_tmp), stride<2>(mSFB_tmp));
|
||||
return make_tensor(mSFB_tmp.data(), make_layout(new_shape, new_stride));
|
||||
}
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
-7
@@ -240,7 +240,6 @@ public:
|
||||
|
||||
// To relax them, we need to handle loading more than 1 row of scales for every main loop iteration.
|
||||
// We must also handle updating the pipeline transaction bytes on the fly.
|
||||
// NOTE: Deleting this assertion without required changes will cause the code to hang.
|
||||
static_assert(size<1>(SmemLayoutAtomScale{}) == 1, "size<1>(SmemLayoutAtomScale) must be 1.");
|
||||
|
||||
private:
|
||||
@@ -490,8 +489,6 @@ public:
|
||||
: args_setup(args.ptr_A, args.ptr_B);
|
||||
}
|
||||
else if constexpr (ModeHasScales) {
|
||||
// NOTE: fix chunk wise scaling
|
||||
//auto scale_k = (K + args.chunk_size - 1) / args.chunk_size;
|
||||
auto scale_k = 1;
|
||||
ElementScale const* ptr_S = reinterpret_cast<ElementScale const*>(args.ptr_S);
|
||||
StrideScale dS{};
|
||||
@@ -998,7 +995,6 @@ public:
|
||||
Utils::copy_tensors_MK(smem_tiled_copy_A, tCsA, tCrA_copy_view,
|
||||
partitioned_extra_info, copy_partitions_extra_info, 0, smem_pipe_read.index());
|
||||
|
||||
// NOTE: Check this when applying swizzling PR on top of GGMD
|
||||
Utils::copy_tensors_MK(smem_tiled_copy_A, tCsA, tCrA_copy_view,
|
||||
partitioned_extra_info, copy_partitions_extra_info, 1, smem_pipe_read.index());
|
||||
|
||||
@@ -1049,7 +1045,6 @@ public:
|
||||
Utils::copy_tensors_MK(smem_tiled_copy_A, tCsA, tCrA_copy_view,
|
||||
partitioned_extra_info, copy_partitions_extra_info, 0, smem_pipe_read.index());
|
||||
|
||||
// NOTE: Check this when applying swizzling PR on top of GGMD
|
||||
Utils::copy_tensors_MK(smem_tiled_copy_A, tCsA, tCrA_copy_view,
|
||||
partitioned_extra_info, copy_partitions_extra_info, 1, smem_pipe_read.index());
|
||||
Utils::dequantize_A_kblock(tCrA_load, tCrA_mma, partitioned_extra_info, 0);
|
||||
@@ -1248,7 +1243,6 @@ public:
|
||||
|
||||
if constexpr (KernelConversionMode == ConversionMode::ConvertAndScale) {
|
||||
NonVoidElementScale const* ptr_S = nullptr;
|
||||
// NOTE: figure out chunk wise scaling. auto scale_k = (K + mainloop_params.chunk_size - 1) / mainloop_params.chunk_size;
|
||||
auto scale_k = 1;
|
||||
Tensor tensor_scale = make_tensor(detail::get_logical_ptr(ptr_S), make_shape(M,scale_k,Int<1>{}), mainloop_params.dS[next_group]);
|
||||
cute::detail::fill_tma_gmem_shape_stride(mainloop_params.tma_load_scale, tensor_scale,
|
||||
@@ -1256,7 +1250,6 @@ public:
|
||||
}
|
||||
else if constexpr (KernelConversionMode == ConversionMode::ConvertAndScaleWithZero) {
|
||||
ElementZero const* ptr_Z = nullptr;
|
||||
// NOTE: figure out chunk wise scaling. auto scale_k = (K + mainloop_params.chunk_size - 1) / mainloop_params.chunk_size;
|
||||
auto scale_k = 1;
|
||||
Tensor tensor_zero = make_tensor(detail::get_logical_ptr(ptr_Z), make_shape(M,scale_k,Int<1>{}), mainloop_params.dS[next_group]);
|
||||
cute::detail::fill_tma_gmem_shape_stride(mainloop_params.tma_load_zero, tensor_zero,
|
||||
|
||||
+1
-2
@@ -531,7 +531,7 @@ struct CollectiveMma<
|
||||
TiledMma tiled_mma;
|
||||
auto thread_mma = tiled_mma.get_slice(warp_group_thread_layout(warp_group_idx));
|
||||
|
||||
Tensor tCsScaleAViewAsC = tiled_mma.get_slice(thread_idx).partition_C(sScaleAViewAsC); // (MMA,MMA_M,MMA_N,PIPE), `thread_mma` above is correct when partitioning A and B, but it is not correct when partitioning C.
|
||||
Tensor tCsScaleAViewAsC = tiled_mma.get_slice(thread_idx).partition_C(sScaleAViewAsC); // (MMA,MMA_M,MMA_N,PIPE), `thread_mma` above is correct when partitioning A and B, but it is not correct when partitioning C.
|
||||
|
||||
Tensor tCsA = thread_mma.partition_A(sA); // (MMA,MMA_M,MMA_K,PIPE)
|
||||
Tensor tCsB = thread_mma.partition_B(sB); // (MMA,MMA_N,MMA_K,PIPE)
|
||||
@@ -557,7 +557,6 @@ struct CollectiveMma<
|
||||
PipelineState smem_pipe_release = smem_pipe_read;
|
||||
|
||||
// Per block scale values for operand A and B
|
||||
|
||||
using RegLayoutScaleAViewAsC = decltype(make_layout_like(tCsScaleAViewAsC(_, _, _, 0).layout())); // `make_layout_like` makes a compact layout.
|
||||
using RegLayoutScaleAEssential = decltype(filter_zeros(RegLayoutScaleAViewAsC{}.stride(), RegLayoutScaleAViewAsC{}.shape())); // an interface to traverse the underlying storage for the compact layout mentioned above
|
||||
|
||||
|
||||
@@ -351,6 +351,23 @@ struct MainloopSm90TmaGmmaWarpSpecializedSparseFP8
|
||||
: MainloopSm90TmaGmmaWarpSpecializedSparse<Stages, ClusterShape, KernelSchedule> {
|
||||
};
|
||||
|
||||
// Mixed precision version n-buffer in rmem (Hopper TMA), pipelined with Hopper GMMA and TMA, Warp specialized dynamic schedule for Ptr-Array and Grouped Gemm
|
||||
template<
|
||||
int Stages_,
|
||||
class ClusterShape_ = Shape<_1,_1,_1>,
|
||||
class KernelSchedule = KernelPtrArrayTmaWarpSpecializedCooperative
|
||||
>
|
||||
struct MainloopSm90ArrayTmaGmmaWarpSpecializedMixedInput {
|
||||
constexpr static int Stages = Stages_;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using ArchTag = arch::Sm90;
|
||||
using Schedule = KernelSchedule;
|
||||
static_assert(
|
||||
cute::is_same_v<Schedule, KernelPtrArrayTmaWarpSpecializedCooperative> ||
|
||||
cute::is_same_v<Schedule, KernelPtrArrayTmaWarpSpecializedPingpong>,
|
||||
"KernelSchedule must be one of the Ptr-Array or Grouped Gemm TMA Warp Specialized Cooperative policies");
|
||||
};
|
||||
|
||||
|
||||
template<
|
||||
int SchedulerPipelineStageCount_,
|
||||
@@ -373,6 +390,16 @@ struct KernelTmaWarpSpecializedBlockScaledSm100 final {
|
||||
|
||||
|
||||
|
||||
// InputTransform GEMM
|
||||
template<
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_
|
||||
>
|
||||
struct KernelTmaWarpSpecializedInputTransformSm100 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_,
|
||||
@@ -393,6 +420,15 @@ struct KernelPtrArrayTmaWarpSpecializedBlockScaledSm100 final {
|
||||
static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_;
|
||||
};
|
||||
|
||||
// Ptr-Array InputTransform GEMM
|
||||
template<
|
||||
int SchedulerPipelineStageCount_,
|
||||
int AccumulatorPipelineStageCount_
|
||||
>
|
||||
struct KernelPtrArrayTmaWarpSpecializedInputTransformSm100 final {
|
||||
static constexpr int SchedulerPipelineStageCount = SchedulerPipelineStageCount_;
|
||||
static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_;
|
||||
};
|
||||
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
@@ -401,32 +437,67 @@ struct KernelPtrArrayTmaWarpSpecializedBlockScaledSm100 final {
|
||||
// Collective Builder Tag Property
|
||||
//
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
//
|
||||
// SM100 Dispatch Policies
|
||||
//
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Base Dispatch Policies
|
||||
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
|
||||
//
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// SM100 Dense GEMM Dispatch Policies
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
struct KernelScheduleSm100DenseGemm : KernelScheduleSm100 {}; // Base policy
|
||||
// Dense GEMM: Specialize for 1SM vs 2SM
|
||||
struct KernelTmaWarpSpecialized1SmSm100 final : KernelSchedule1Sm, KernelScheduleSm100DenseGemm {};
|
||||
struct KernelTmaWarpSpecialized2SmSm100 final : KernelSchedule2Sm, KernelScheduleSm100DenseGemm {};
|
||||
struct KernelTmaWarpSpecialized1SmSm100 final : KernelSchedule1Sm, KernelScheduleSm100DenseGemm {}; // Use for 1SM Dense GEMM Kernels for Collective Mainloop Builder
|
||||
struct KernelTmaWarpSpecialized2SmSm100 final : KernelSchedule2Sm, KernelScheduleSm100DenseGemm {}; // Use for 2SM Dense GEMM Kernels for Collective Mainloop Builder
|
||||
// Dense GEMM + (Ptr Array or Group GEMM)
|
||||
struct KernelScheduleSm100PtrArrayDenseGemm : KernelScheduleSm100DenseGemm {};
|
||||
// Ptr-Array Dense GEMM: Specialize for 1SM vs 2SM
|
||||
struct KernelPtrArrayTmaWarpSpecialized1SmSm100 final : KernelSchedule1Sm, KernelScheduleSm100PtrArrayDenseGemm {};
|
||||
struct KernelPtrArrayTmaWarpSpecialized2SmSm100 final : KernelSchedule2Sm, KernelScheduleSm100PtrArrayDenseGemm {};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// SM100 Planar Complex GEMM Dispatch Policies
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
struct KernelScheduleSm100PlanarComplexGemm : KernelScheduleSm100{};
|
||||
// Planar Complex GEMM: Specialize for 1SM vs 2SM
|
||||
struct KernelTmaWarpSpecialized1SmPlanarComplexSm100 final : KernelSchedule1Sm, KernelScheduleSm100PlanarComplexGemm { };
|
||||
struct KernelTmaWarpSpecialized2SmPlanarComplexSm100 final : KernelSchedule2Sm, KernelScheduleSm100PlanarComplexGemm { };
|
||||
// Planar Complex GEMM + (Ptr Array or Group GEMM)
|
||||
struct KernelScheduleSm100PtrArrayPlanarComplexGemm : KernelScheduleSm100PlanarComplexGemm {};
|
||||
struct KernelPtrArrayTmaWarpSpecialized1SmPlanarComplexSm100 final : KernelSchedule1Sm, KernelScheduleSm100PtrArrayPlanarComplexGemm {};
|
||||
struct KernelPtrArrayTmaWarpSpecialized2SmPlanarComplexSm100 final : KernelSchedule2Sm, KernelScheduleSm100PtrArrayPlanarComplexGemm {};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// SM100 FastF32 (9xBF16) GEMM Dispatch Policies
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
struct KernelScheduleSm100FastFP32Gemm : KernelScheduleSm100 {};
|
||||
struct KernelTmaWarpSpecializedFastFP32SmemSm100 : KernelScheduleSm100FastFP32Gemm { };
|
||||
// Dispatch policies without smem load the A operand from tmem
|
||||
struct KernelTmaWarpSpecialized1SmFastFP32Sm100 final : KernelSchedule1Sm, KernelScheduleSm100FastFP32Gemm { };
|
||||
struct KernelTmaWarpSpecialized2SmFastFP32Sm100 final : KernelSchedule2Sm, KernelScheduleSm100FastFP32Gemm { };
|
||||
// Dispatch policies with smem load the A operand from smem
|
||||
struct KernelTmaWarpSpecialized1SmFastFP32SmemSm100 final : KernelSchedule1Sm, KernelTmaWarpSpecializedFastFP32SmemSm100 { };
|
||||
struct KernelTmaWarpSpecialized2SmFastFP32SmemSm100 final : KernelSchedule2Sm, KernelTmaWarpSpecializedFastFP32SmemSm100 { };
|
||||
// Ptr-Array Transform GEMM: Specialize for 1SM vs 2SM FastF32 GEMM
|
||||
struct KernelScheduleSm100PtrArrayFastFP32Gemm : KernelScheduleSm100FastFP32Gemm {};
|
||||
struct KernelTmaWarpSpecializedPtrArrayFastFP32SmemSm100 : KernelScheduleSm100PtrArrayFastFP32Gemm { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized1SmFastFP32Sm100 final : KernelSchedule1Sm, KernelScheduleSm100PtrArrayFastFP32Gemm { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized2SmFastFP32Sm100 final : KernelSchedule2Sm, KernelScheduleSm100PtrArrayFastFP32Gemm { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized1SmFastFP32SmemSm100 final : KernelSchedule1Sm, KernelTmaWarpSpecializedPtrArrayFastFP32SmemSm100 { };
|
||||
struct KernelPtrArrayTmaWarpSpecialized2SmFastFP32SmemSm100 final : KernelSchedule2Sm, KernelTmaWarpSpecializedPtrArrayFastFP32SmemSm100 { };
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
// SM100 BlockScaled Dense GEMM Dispatch Policies
|
||||
///////////////////////////////////////////////////////////////////////////////////////////////////////
|
||||
struct KernelScheduleBlockScaledGemmSm100 : KernelScheduleSm100 {};
|
||||
struct KernelScheduleMxNvf4Sm100 : KernelScheduleBlockScaledGemmSm100 {};
|
||||
struct KernelScheduleMxf8f6f4Sm100 : KernelScheduleBlockScaledGemmSm100 {};
|
||||
// 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 { };
|
||||
@@ -436,13 +507,10 @@ struct KernelTmaWarpSpecialized1SmMxf4Sm100 final : KernelSchedule1Sm, KernelSch
|
||||
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 {};
|
||||
|
||||
|
||||
// BlockScaled Dense GEMM + (Ptr Array or Group GEMM)
|
||||
struct KernelSchedulePtrArrayBlockScaledGemmSm100 : KernelScheduleBlockScaledGemmSm100 {};
|
||||
struct KernelSchedulePtrArrayMxNvf4Sm100 : KernelSchedulePtrArrayBlockScaledGemmSm100 {};
|
||||
struct KernelSchedulePtrArrayMxf8f6f4Sm100 : KernelSchedulePtrArrayBlockScaledGemmSm100 {};
|
||||
// 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 { };
|
||||
@@ -454,6 +522,7 @@ struct KernelPtrArrayTmaWarpSpecialized1SmMxf8f6f4Sm100 final : KernelSchedule1S
|
||||
struct KernelPtrArrayTmaWarpSpecialized2SmMxf8f6f4Sm100 final : KernelSchedule2Sm, KernelSchedulePtrArrayMxf8f6f4Sm100 { };
|
||||
|
||||
|
||||
|
||||
// n-buffer in smem, pipelined with Blackwell UMMA and TMA, Warp specialized dynamic schedule
|
||||
template<
|
||||
int Stages_,
|
||||
@@ -488,6 +557,55 @@ struct MainloopSm100TmaUmmaWarpSpecializedBlockScaled {
|
||||
|
||||
|
||||
|
||||
// n-buffer in smem, pipelined with Blackwell Fast FP32 kernel with UMMA (HwScaled) and TMA,
|
||||
// Warp specialized dynamic schedule
|
||||
template<
|
||||
// Number of Pipeline stages for
|
||||
// MainloopLoad <-> Conversion <-> MainLoad
|
||||
int Load2TransformPipelineStageCount_,
|
||||
// Number of Pipeline stages for
|
||||
// MainloopLoad <-> Conversion <-> MainLoad
|
||||
int Transform2MmaPipelineStageCount_,
|
||||
// TileScheduler pipeline depth
|
||||
int SchedulerPipelineStageCount_,
|
||||
// Accmulator pipeline depth
|
||||
int AccumulatorPipelineStageCount_,
|
||||
// Number of MMA Bands to be computed in a single FastF32 MMA operation.
|
||||
// For BF16 emulation, we have 3 compute matrices, with 9 MMAs forming 5 bands.
|
||||
// We can eliminate bands 4 and/or 5 (up to last 3 MMA operations).
|
||||
// Valid values are 3, 4, 5
|
||||
int NumBandsToCompute_,
|
||||
// Scaling factor for decomposed matrices (2^ScalingFactor)
|
||||
// 8 for BF16, 11 for TF32
|
||||
int ScalingFactor_,
|
||||
// Number of UMMA instructions emulated a single stage
|
||||
// Ex: Staged16 has 1 FastF32 MMA per stage
|
||||
// Should be smaller than K-mode of a single ClusterTile
|
||||
int AccPromotionInterval_,
|
||||
// ClusterShape for the kernel
|
||||
class ClusterShape_ = Shape<_1,_1,_1>,
|
||||
// The TMEM_LOAD atom to be used for loading local accumulator
|
||||
// from TMEM to registers
|
||||
class AccumulatorCopyAtom_ = cute::SM100_TMEM_LOAD_32dp32b32x
|
||||
>
|
||||
struct MainloopSm100TmaUmmaWarpSpecializedFastF32 {
|
||||
constexpr static int Load2TransformPipelineStageCount = Load2TransformPipelineStageCount_;
|
||||
constexpr static int Transform2MmaPipelineStageCount = Transform2MmaPipelineStageCount_;
|
||||
constexpr static int NumBandsToCompute = NumBandsToCompute_;
|
||||
constexpr static int ScalingFactor = ScalingFactor_;
|
||||
constexpr static int AccPromotionInterval = AccPromotionInterval_;
|
||||
constexpr static detail::KernelInputTransformType InputTransformType = detail::KernelInputTransformType::FastF32;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using AccumulatorCopyAtom = AccumulatorCopyAtom_;
|
||||
using ArchTag = arch::Sm100;
|
||||
using Schedule = KernelTmaWarpSpecializedInputTransformSm100<SchedulerPipelineStageCount_, AccumulatorPipelineStageCount_>;
|
||||
|
||||
// For backwards compatibility with GemmUniversalAdapter.
|
||||
constexpr static int Stages = Load2TransformPipelineStageCount;
|
||||
};
|
||||
|
||||
|
||||
|
||||
// n-buffer in smem, pipelined with Blackwell UMMA and TMA, Warp specialized dynamic schedule
|
||||
template<
|
||||
int Stages_,
|
||||
@@ -520,6 +638,55 @@ struct MainloopSm100ArrayTmaUmmaWarpSpecializedBlockScaled {
|
||||
|
||||
|
||||
|
||||
// n-buffer in smem, pipelined with Blackwell Fast FP32 kernel with UMMA (HwScaled) and TMA,
|
||||
// Warp specialized dynamic schedule
|
||||
template<
|
||||
// Number of Pipeline stages for
|
||||
// MainloopLoad <-> Conversion <-> MainLoad
|
||||
int Load2TransformPipelineStageCount_,
|
||||
// Number of Pipeline stages for
|
||||
// MainloopLoad <-> Conversion <-> MainLoad
|
||||
int Transform2MmaPipelineStageCount_,
|
||||
// TileScheduler pipeline depth
|
||||
int SchedulerPipelineStageCount_,
|
||||
// Accmulator pipeline depth
|
||||
int AccumulatorPipelineStageCount_,
|
||||
// Number of MMA Bands to be computed in a single FastF32 MMA operation.
|
||||
// For BF16 emulation, we have 3 compute matrices, with 9 MMAs forming 5 bands.
|
||||
// We can eliminate bands 4 and/or 5 (up to last 3 MMA operations).
|
||||
// Valid values are 3, 4, 5
|
||||
int NumBandsToCompute_,
|
||||
// Scaling factor for decomposed matrices (2^ScalingFactor)
|
||||
// 8 for BF16, 11 for TF32
|
||||
int ScalingFactor_,
|
||||
// Number of UMMA instructions emulated a single stage
|
||||
// Ex: Staged16 has 1 FastF32 MMA per stage
|
||||
// Should be smaller than K-mode of a single ClusterTile
|
||||
int AccPromotionInterval_,
|
||||
// ClusterShape for the kernel
|
||||
class ClusterShape_ = Shape<_1,_1,_1>,
|
||||
// The TMEM_LOAD atom to be used for loading local accumulator
|
||||
// from TMEM to registers
|
||||
class AccumulatorCopyAtom_ = cute::SM100_TMEM_LOAD_32dp32b32x
|
||||
>
|
||||
struct MainloopSm100ArrayTmaUmmaWarpSpecializedFastF32 {
|
||||
constexpr static int Load2TransformPipelineStageCount = Load2TransformPipelineStageCount_;
|
||||
constexpr static int Transform2MmaPipelineStageCount = Transform2MmaPipelineStageCount_;
|
||||
constexpr static int NumBandsToCompute = NumBandsToCompute_;
|
||||
constexpr static int ScalingFactor = ScalingFactor_;
|
||||
constexpr static int AccPromotionInterval = AccPromotionInterval_;
|
||||
constexpr static detail::KernelInputTransformType InputTransformType = detail::KernelInputTransformType::FastF32;
|
||||
using ClusterShape = ClusterShape_;
|
||||
using AccumulatorCopyAtom = AccumulatorCopyAtom_;
|
||||
using ArchTag = arch::Sm100;
|
||||
using Schedule = KernelPtrArrayTmaWarpSpecializedInputTransformSm100<SchedulerPipelineStageCount_, AccumulatorPipelineStageCount_>;
|
||||
|
||||
// For backwards compatibility with GemmUniversalAdapter.
|
||||
constexpr static int Stages = Load2TransformPipelineStageCount;
|
||||
};
|
||||
|
||||
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
} // namespace cutlass::gemm
|
||||
|
||||
@@ -63,6 +63,10 @@ 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"
|
||||
|
||||
#include "cutlass/gemm/kernel/sm100_gemm_tma_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/kernel/sm100_gemm_array_tma_warpspecialized.hpp"
|
||||
#include "cutlass/gemm/kernel/sm100_gemm_tma_warpspecialized_input_transform.hpp"
|
||||
#include "cutlass/gemm/kernel/sm100_gemm_array_tma_warpspecialized_input_transform.hpp"
|
||||
|
||||
////////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
+1139
File diff suppressed because it is too large
Load Diff
File diff suppressed because it is too large
Load Diff
@@ -120,7 +120,7 @@ public:
|
||||
typename detail::TileSchedulerSelector<
|
||||
GroupScheduler, ArchTag,
|
||||
TileShape, ClusterShape,
|
||||
2, // Default unused parameter - SchedulerPipelineStageCoun
|
||||
2, // Default unused parameter - SchedulerPipelineStageCount
|
||||
ProblemShape>::Scheduler,
|
||||
typename detail::TileSchedulerSelector<
|
||||
void, ArchTag, TileShape, ClusterShape>::Scheduler>;
|
||||
|
||||
@@ -120,7 +120,7 @@ public:
|
||||
typename detail::TileSchedulerSelector<
|
||||
GroupScheduler, ArchTag,
|
||||
TileShape, ClusterShape,
|
||||
2, // Default unused parameter - SchedulerPipelineStageCoun
|
||||
2, // Default unused parameter - SchedulerPipelineStageCount
|
||||
ProblemShape>::Scheduler,
|
||||
typename detail::TileSchedulerSelector<
|
||||
void, ArchTag, TileShape, ClusterShape>::Scheduler>;
|
||||
|
||||
Reference in New Issue
Block a user