v4.5 tag update (#3202)

* Python DSL examples reorganization.

* v4.5 tag update.
This commit is contained in:
Junkai-Wu
2026-05-05 20:55:27 -04:00
committed by GitHub
parent f74fea9ce3
commit cb37157db5
351 changed files with 36688 additions and 8117 deletions
+5
View File
@@ -209,3 +209,8 @@
# define CUTE_ARCH_TCGEN05_MXF4NVF4_MMA_ULTRA_ENABLED
#endif
#if (defined(CUTLASS_ARCH_MMA_SM103A_ENABLED) || defined(CUTLASS_ARCH_MMA_SM103F_ENABLED))
#define CUTE_ARCH_TCGEN05_TMEM_STAT_ENABLED
#endif
+169
View File
@@ -7612,4 +7612,173 @@ struct SM100_TMEM_STORE_32dp32b128x_16b
////////////////////////////////////////////////////////////////////////////////////////////////////
////////////////////////////////////////////////////////////////////////////////////////////////////
namespace SM100::TMEM::LOAD_STAT {
// 32 data path lanes, 32-bit pattern, repeated 32 times
struct SM100_TMEM_LOAD_STAT_32dp32b32x
{
using SRegisters = uint32_t[1];
using DRegisters = uint32_t[32];
CUTE_HOST_DEVICE static void
copy(uint32_t const& src_addr,
uint32_t& dst00, uint32_t& dst01, uint32_t& dst02, uint32_t& dst03,
uint32_t& dst04, uint32_t& dst05, uint32_t& dst06, uint32_t& dst07,
uint32_t& dst08, uint32_t& dst09, uint32_t& dst10, uint32_t& dst11,
uint32_t& dst12, uint32_t& dst13, uint32_t& dst14, uint32_t& dst15,
uint32_t& dst16, uint32_t& dst17, uint32_t& dst18, uint32_t& dst19,
uint32_t& dst20, uint32_t& dst21, uint32_t& dst22, uint32_t& dst23,
uint32_t& dst24, uint32_t& dst25, uint32_t& dst26, uint32_t& dst27,
uint32_t& dst28, uint32_t& dst29, uint32_t& dst30, uint32_t& dst31,
float& row_max)
{
#if defined(CUTE_ARCH_TCGEN05_TMEM_STAT_ENABLED)
asm volatile ("tcgen05.ld.red.sync.aligned.32x32b.x32.max.f32"
"{%0, %1, %2, %3,"
"%4, %5, %6, %7,"
"%8, %9, %10, %11,"
"%12, %13, %14, %15,"
"%16, %17, %18, %19,"
"%20, %21, %22, %23,"
"%24, %25, %26, %27,"
"%28, %29, %30, %31}, %32,"
"[%33];\n"
: "=r"(dst00), "=r"(dst01), "=r"(dst02), "=r"(dst03),
"=r"(dst04), "=r"(dst05), "=r"(dst06), "=r"(dst07),
"=r"(dst08), "=r"(dst09), "=r"(dst10), "=r"(dst11),
"=r"(dst12), "=r"(dst13), "=r"(dst14), "=r"(dst15),
"=r"(dst16), "=r"(dst17), "=r"(dst18), "=r"(dst19),
"=r"(dst20), "=r"(dst21), "=r"(dst22), "=r"(dst23),
"=r"(dst24), "=r"(dst25), "=r"(dst26), "=r"(dst27),
"=r"(dst28), "=r"(dst29), "=r"(dst30), "=r"(dst31),
"=f"(row_max)
: "r"(src_addr));
#else
CUTE_INVALID_CONTROL_PATH("Trying to use TMEM_LOAD_STAT without CUTE_ARCH_TCGEN05_TMEM_STAT_ENABLED.");
#endif
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
// 32 data path lanes, 32-bit pattern, repeated 128 times
struct SM100_TMEM_LOAD_STAT_32dp32b128x
{
using SRegisters = uint32_t[1];
using DRegisters = uint32_t[128];
CUTE_HOST_DEVICE static void
copy(uint32_t const& src_addr,
uint32_t& dst000, uint32_t& dst001, uint32_t& dst002, uint32_t& dst003,
uint32_t& dst004, uint32_t& dst005, uint32_t& dst006, uint32_t& dst007,
uint32_t& dst008, uint32_t& dst009, uint32_t& dst010, uint32_t& dst011,
uint32_t& dst012, uint32_t& dst013, uint32_t& dst014, uint32_t& dst015,
uint32_t& dst016, uint32_t& dst017, uint32_t& dst018, uint32_t& dst019,
uint32_t& dst020, uint32_t& dst021, uint32_t& dst022, uint32_t& dst023,
uint32_t& dst024, uint32_t& dst025, uint32_t& dst026, uint32_t& dst027,
uint32_t& dst028, uint32_t& dst029, uint32_t& dst030, uint32_t& dst031,
uint32_t& dst032, uint32_t& dst033, uint32_t& dst034, uint32_t& dst035,
uint32_t& dst036, uint32_t& dst037, uint32_t& dst038, uint32_t& dst039,
uint32_t& dst040, uint32_t& dst041, uint32_t& dst042, uint32_t& dst043,
uint32_t& dst044, uint32_t& dst045, uint32_t& dst046, uint32_t& dst047,
uint32_t& dst048, uint32_t& dst049, uint32_t& dst050, uint32_t& dst051,
uint32_t& dst052, uint32_t& dst053, uint32_t& dst054, uint32_t& dst055,
uint32_t& dst056, uint32_t& dst057, uint32_t& dst058, uint32_t& dst059,
uint32_t& dst060, uint32_t& dst061, uint32_t& dst062, uint32_t& dst063,
uint32_t& dst064, uint32_t& dst065, uint32_t& dst066, uint32_t& dst067,
uint32_t& dst068, uint32_t& dst069, uint32_t& dst070, uint32_t& dst071,
uint32_t& dst072, uint32_t& dst073, uint32_t& dst074, uint32_t& dst075,
uint32_t& dst076, uint32_t& dst077, uint32_t& dst078, uint32_t& dst079,
uint32_t& dst080, uint32_t& dst081, uint32_t& dst082, uint32_t& dst083,
uint32_t& dst084, uint32_t& dst085, uint32_t& dst086, uint32_t& dst087,
uint32_t& dst088, uint32_t& dst089, uint32_t& dst090, uint32_t& dst091,
uint32_t& dst092, uint32_t& dst093, uint32_t& dst094, uint32_t& dst095,
uint32_t& dst096, uint32_t& dst097, uint32_t& dst098, uint32_t& dst099,
uint32_t& dst100, uint32_t& dst101, uint32_t& dst102, uint32_t& dst103,
uint32_t& dst104, uint32_t& dst105, uint32_t& dst106, uint32_t& dst107,
uint32_t& dst108, uint32_t& dst109, uint32_t& dst110, uint32_t& dst111,
uint32_t& dst112, uint32_t& dst113, uint32_t& dst114, uint32_t& dst115,
uint32_t& dst116, uint32_t& dst117, uint32_t& dst118, uint32_t& dst119,
uint32_t& dst120, uint32_t& dst121, uint32_t& dst122, uint32_t& dst123,
uint32_t& dst124, uint32_t& dst125, uint32_t& dst126, uint32_t& dst127, float& row_max)
{
#if defined(CUTE_ARCH_TCGEN05_TMEM_STAT_ENABLED)
asm volatile ("tcgen05.ld.red.sync.aligned.32x32b.x128.max.f32"
"{%0, %1, %2, %3,"
"%4, %5, %6, %7,"
"%8, %9, %10, %11,"
"%12, %13, %14, %15,"
"%16, %17, %18, %19,"
"%20, %21, %22, %23,"
"%24, %25, %26, %27,"
"%28, %29, %30, %31,"
"%32, %33, %34, %35,"
"%36, %37, %38, %39,"
"%40, %41, %42, %43,"
"%44, %45, %46, %47,"
"%48, %49, %50, %51,"
"%52, %53, %54, %55,"
"%56, %57, %58, %59,"
"%60, %61, %62, %63,"
"%64, %65, %66, %67,"
"%68, %69, %70, %71,"
"%72, %73, %74, %75,"
"%76, %77, %78, %79,"
"%80, %81, %82, %83,"
"%84, %85, %86, %87,"
"%88, %89, %90, %91,"
"%92, %93, %94, %95,"
"%96, %97, %98, %99,"
"%100, %101, %102, %103,"
"%104, %105, %106, %107,"
"%108, %109, %110, %111,"
"%112, %113, %114, %115,"
"%116, %117, %118, %119,"
"%120, %121, %122, %123,"
"%124, %125, %126, %127}, %128,"
"[%129];\n"
: "=r"(dst000), "=r"(dst001), "=r"(dst002), "=r"(dst003),
"=r"(dst004), "=r"(dst005), "=r"(dst006), "=r"(dst007),
"=r"(dst008), "=r"(dst009), "=r"(dst010), "=r"(dst011),
"=r"(dst012), "=r"(dst013), "=r"(dst014), "=r"(dst015),
"=r"(dst016), "=r"(dst017), "=r"(dst018), "=r"(dst019),
"=r"(dst020), "=r"(dst021), "=r"(dst022), "=r"(dst023),
"=r"(dst024), "=r"(dst025), "=r"(dst026), "=r"(dst027),
"=r"(dst028), "=r"(dst029), "=r"(dst030), "=r"(dst031),
"=r"(dst032), "=r"(dst033), "=r"(dst034), "=r"(dst035),
"=r"(dst036), "=r"(dst037), "=r"(dst038), "=r"(dst039),
"=r"(dst040), "=r"(dst041), "=r"(dst042), "=r"(dst043),
"=r"(dst044), "=r"(dst045), "=r"(dst046), "=r"(dst047),
"=r"(dst048), "=r"(dst049), "=r"(dst050), "=r"(dst051),
"=r"(dst052), "=r"(dst053), "=r"(dst054), "=r"(dst055),
"=r"(dst056), "=r"(dst057), "=r"(dst058), "=r"(dst059),
"=r"(dst060), "=r"(dst061), "=r"(dst062), "=r"(dst063),
"=r"(dst064), "=r"(dst065), "=r"(dst066), "=r"(dst067),
"=r"(dst068), "=r"(dst069), "=r"(dst070), "=r"(dst071),
"=r"(dst072), "=r"(dst073), "=r"(dst074), "=r"(dst075),
"=r"(dst076), "=r"(dst077), "=r"(dst078), "=r"(dst079),
"=r"(dst080), "=r"(dst081), "=r"(dst082), "=r"(dst083),
"=r"(dst084), "=r"(dst085), "=r"(dst086), "=r"(dst087),
"=r"(dst088), "=r"(dst089), "=r"(dst090), "=r"(dst091),
"=r"(dst092), "=r"(dst093), "=r"(dst094), "=r"(dst095),
"=r"(dst096), "=r"(dst097), "=r"(dst098), "=r"(dst099),
"=r"(dst100), "=r"(dst101), "=r"(dst102), "=r"(dst103),
"=r"(dst104), "=r"(dst105), "=r"(dst106), "=r"(dst107),
"=r"(dst108), "=r"(dst109), "=r"(dst110), "=r"(dst111),
"=r"(dst112), "=r"(dst113), "=r"(dst114), "=r"(dst115),
"=r"(dst116), "=r"(dst117), "=r"(dst118), "=r"(dst119),
"=r"(dst120), "=r"(dst121), "=r"(dst122), "=r"(dst123),
"=r"(dst124), "=r"(dst125), "=r"(dst126), "=r"(dst127), "=f"(row_max)
: "r"(src_addr));
#else
CUTE_INVALID_CONTROL_PATH("Trying to use TMEM_LOAD_STAT without CUTE_ARCH_TCGEN05_TMEM_STAT_ENABLED.");
#endif
}
//#endif
};
} // end namespace SM100::TMEM::LOAD_STAT
} // end namespace cute
+99
View File
@@ -3506,6 +3506,105 @@ tmem_load_to_store(CopyOp) {
////////////////////////////////////////////////////////////////////////////////////////////////////
namespace SM100::TMEM::LOAD_STAT {
//
// Specialized copy_unpack implementation for SM100::TMEM::LOAD_STAT instructions
//
template <class CopyOp,
class TS, class SLayout,
class TD, class DLayout>
CUTE_HOST_DEVICE constexpr
void
copy_unpack(Copy_Traits<CopyOp> const& traits,
Tensor<TS,SLayout> const& src,
Tensor<TD,DLayout> & dst)
{
static_assert(is_tmem<TS>::value, "Expected TMEM src.");
static_assert(is_rmem<TD>::value, "Expected RMEM dst.");
using SrcType = typename TS::value_type;
CUTE_STATIC_ASSERT_V((coalesce(layout(src)) == coalesce(upcast<sizeof_bits<SrcType>::value>(typename Copy_Traits<CopyOp>::ValID{}))),
"Expected src to have the specific TMEM layout required by CopyOp.");
uint32_t tmem_addr = raw_pointer_cast(src.data());
const float& row_max = traits.get_max();
using RegTypeDst = typename remove_extent<typename CopyOp::DRegisters>::type;
Tensor rD = recast<RegTypeDst>(dst);
constexpr int RegNumDst = extent<typename CopyOp::DRegisters>::value;
CUTE_STATIC_ASSERT_V(size(rD) == Int<RegNumDst>{},
"In CopyAtom, dst layout doesn't vectorize into registers. This dst layout is incompatible with this CopyOp.");
// thread idx <=> DP lane assert.
// ASSERT thread attemping to access DP lane within sub-partition.
#if defined(__CUDA_ARCH__) && !defined(NDEBUG)
assert(((uint32_t(threadIdx.x) / 32) % 4) == (((tmem_addr >> 16) / 32) % 4));
#endif
float* row_max_ = const_cast<float*>(&row_max);
float tmp_row_max = row_max_[0];
detail::explode(CopyOp::copy,
&tmem_addr, seq<0>{},
rD, make_seq<RegNumDst>{},
&tmp_row_max, seq<0>{});
row_max_[0] = fmax(row_max_[0], tmp_row_max);
}
} // end namespace SM100::TMEM::LOAD_STAT
////////////////////////////////////////////////////////////////////////////////////////////////////
using SM100::TMEM::LOAD_STAT::SM100_TMEM_LOAD_STAT_32dp32b32x;
template <>
struct Copy_Traits<SM100_TMEM_LOAD_STAT_32dp32b32x>
{
using ThrID = Layout<_32>;
using ValID = Layout<Shape <_1024, Int<32>>,
Stride< _1,TMEM::DP_b>>;
using SrcLayout = Layout<Shape <_32,_32768>,
Stride< _0, _1>>;
using DstLayout = Layout<Shape < _32, Int<1024>>,
Stride<Int<1024>, _1>>;
using RefLayout = SrcLayout;
float row_max = -cutlass::platform::numeric_limits<float>::infinity();
CUTE_HOST_DEVICE constexpr
float const& get_max() const {
return row_max;
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
using SM100::TMEM::LOAD_STAT::SM100_TMEM_LOAD_STAT_32dp32b128x;
template <>
struct Copy_Traits<SM100_TMEM_LOAD_STAT_32dp32b128x>
{
using ThrID = Layout<_32>;
using ValID = Layout<Shape <_4096, _32>,
Stride< _1,TMEM::DP_b>>;
using SrcLayout = Layout<Shape <_32,_131072>,
Stride< _0, _1>>;
using DstLayout = Layout<Shape < _32,_4096>,
Stride<_4096, _1>>;
using RefLayout = SrcLayout;
float row_max = -cutlass::platform::numeric_limits<float>::infinity();
CUTE_HOST_DEVICE constexpr
float const& get_max() const {
return row_max;
}
};
////////////////////////////////////////////////////////////////////////////////////////////////////
////////////////////////////////////////////////////////////////////////////////////////////////////
//
// UTCCP Copy Traits
+1 -5
View File
@@ -333,11 +333,7 @@ struct DescriptorIterator
CUTE_HOST_DEVICE constexpr
DescriptorIterator operator+(Index const& offset) const
{
// Use 32bit calculation rather than 64 bit calculation as we only update the part of desc
SmemDescriptor ret;
ret.lo = desc_.lo + uint32_t(offset);
ret.hi = desc_.hi;
return { ret };
return { desc_ + uint64_t(offset)};
}
};
+2
View File
@@ -49,8 +49,10 @@
#endif
#if ((__CUDACC_VER_MAJOR__ >= 12) || ((__CUDACC_VER_MAJOR__ == 11) && (__CUDACC_VER_MINOR__ >= 8)))
#if !(defined(__QNX__) && __QNX__ >= 800 && defined(NV_IS_SAFETY))
# define CUTLASS_SM90_CLUSTER_LAUNCH_ENABLED
#endif
#endif
#if (__CUDACC_VER_MAJOR__ > 12 || (__CUDACC_VER_MAJOR__ == 12 && __CUDACC_VER_MINOR__ >= 8))
# define CUDA_ENABLE_PREFERRED_CLUSTER
@@ -269,7 +269,7 @@ public:
bool implementable = true;
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if constexpr (IsDynamicCluster) {
static constexpr int MaxClusterSize = 16;
+2
View File
@@ -87,7 +87,9 @@ namespace cutlass {
#if ((__CUDACC_VER_MAJOR__ >= 12) || \
((__CUDACC_VER_MAJOR__ == 11) && (__CUDACC_VER_MINOR__ >= 8)))
#if !(defined(__QNX__) && __QNX__ >= 800 && defined(NV_IS_SAFETY))
#include <cudaTypedefs.h>
#endif
#endif // (__CUDACC_VERSION__ >= 11.8)
#include <driver_types.h>
@@ -121,7 +121,8 @@ struct CollectiveBuilder<
StageCountType,
BuilderScheduleTag,
cute::enable_if_t<
cute::is_same_v<KernelMixedTmaCpAsyncWarpSpecialized1SmBlockScaledSm100, BuilderScheduleTag> &&
(cute::is_same_v<KernelMixedTmaCpAsyncWarpSpecialized1SmBlockScaledSm100, BuilderScheduleTag> ||
cute::is_same_v<KernelMixedTmaCpAsyncWarpSpecialized2SmBlockScaledSm100, BuilderScheduleTag>) &&
(cute::is_same_v<ArchTag, arch::Sm100>
)
>
@@ -139,7 +140,7 @@ struct CollectiveBuilder<
static_assert(cute::is_static_v<TileShape_MNK>, "TileShape has to be static");
static_assert(detail::blockscaled::check_input_datatypes<BuilderScheduleTag, ElementPairA, ElementPairB, UmmaMajorA, UmmaMajorB>(), "Incorrect input types");
static constexpr bool is_2sm = false; // detail::blockscaled::is_2sm<TileShape_MNK, ClusterShape_MNK, BuilderScheduleTag>();
static constexpr bool is_2sm = detail::blockscaled::is_2sm<TileShape_MNK, ClusterShape_MNK, BuilderScheduleTag>();
static constexpr auto Instr = detail::blockscaled::select_instr<ElementPairA, ElementPairB, ElementAccumulator, UmmaMajorA, UmmaMajorB, BuilderScheduleTag>();
using TiledMma = typename cutlass::gemm::collective::detail::TrivialBlockscaledMma<ElementPairA, ElementPairB, ElementAccumulator,
@@ -180,7 +181,7 @@ struct CollectiveBuilder<
cute::size<2>(TileShape_MNK{}))));
// Assigning 4 warps for mainloop load of B
static constexpr int NumLoadThreadsCpAsync = 128;
static constexpr int NumLoadThreadsCpAsync = 128 / size(AtomThrID{});
using SmemShapeA_M = decltype(shape_div(shape<0>(TileShape_MNK{}), shape_div(shape<0>(TileShape_MNK{}), size<0>(TileShape_MNK{}) / size(AtomThrID{}))));
@@ -196,7 +197,7 @@ struct CollectiveBuilder<
using GmemCopyAtomB = cute::Copy_Atom<SM80_CP_ASYNC_CACHEGLOBAL_ZFILL<AlignmentTypeB>, ElementB>;
using GmemTiledCopyB = decltype(detail::make_simt_gmem_tiled_copy<
GmemCopyAtomB, NumLoadThreadsCpAsync, AlignmentB, TagToStrideB_t<GmemLayoutBTag>,
decltype(cute::get<1>(TileShape_MNK{})), decltype(cute::get<2>(TileShape_MNK{}))>());
decltype(cute::get<1>(TileShape_MNK{}) / size(AtomThrID{})), decltype(cute::get<2>(TileShape_MNK{}))>());
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{}));
@@ -233,15 +234,24 @@ struct CollectiveBuilder<
static constexpr uint32_t SchedulerPipelineStageCount = AccumulatorPipelineStageCount + 1;
// AccumulatorPipeline = PipelineUmmaAsync
static constexpr auto AccumulatorPipelineStorage = sizeof(typename cutlass::PipelineUmmaAsync<AccumulatorPipelineStageCount>::SharedStorage);
static constexpr auto AccumulatorPipelineStorage = sizeof(typename cutlass::PipelineUmmaAsync<AccumulatorPipelineStageCount, AtomThrID>::SharedStorage);
// CLCPipeline = PipelineCLCFetchAsync
static constexpr auto CLCPipelineStorage = sizeof(typename cutlass::PipelineCLCFetchAsync<SchedulerPipelineStageCount, ClusterShape_MNK>::SharedStorage);
// CLC (scheduler) response
static constexpr auto CLCResponseStorage = SchedulerPipelineStageCount * detail::CLCResponseSize;
// Tmem dealloc barrier
static constexpr auto TmemDeallocStorage = sizeof(cutlass::arch::ClusterBarrier);
// MMA trampoline barrier (for 2SM synchronization)
static constexpr auto MmaTrampolineBarrierStorage = sizeof(cutlass::arch::ClusterBarrier);
// Tmem base pointer storage
static constexpr auto TmemBasePtrStorage = sizeof(uint32_t);
// Smem usage that's not part of CollectiveEpilogue::SharedStorage & CollectiveMainloop::SharedStorage
static constexpr auto KernelSmemCarveout = static_cast<int>( AccumulatorPipelineStorage +
CLCPipelineStorage +
CLCResponseStorage);
CLCResponseStorage +
TmemDeallocStorage +
MmaTrampolineBarrierStorage +
TmemBasePtrStorage);
// Reduce SMEM capacity available for buffers considering barrier allocations.
static constexpr int ReducedSmemCapacityBytes = detail::sm100_reduced_smem_capacity_bytes<ArchTag, KernelSmemCarveout>();
@@ -124,6 +124,7 @@ struct CollectiveBuilder<
) &&
// Blockscaled Gemm
(not cute::is_same_v<KernelMixedTmaCpAsyncWarpSpecialized1SmBlockScaledSm100, BuilderScheduleTag>) &&
(not cute::is_same_v<KernelMixedTmaCpAsyncWarpSpecialized2SmBlockScaledSm100, BuilderScheduleTag>) &&
(cute::is_base_of_v<KernelScheduleBlockScaledGemmSm100, BuilderScheduleTag> ||
cute::is_same_v<KernelScheduleAuto, BuilderScheduleTag>)
&&
@@ -65,10 +65,11 @@ struct CollectiveBuilder<
AlignmentB,
ElementAccumulator,
TileShape_MNK, // (MmaAtomShapeM, MmaAtomShapeN, TileK)
ClusterShape_MNK, // Static cluster shape (_1, _1, _1)
ClusterShape_MNK, // Static cluster shape, but can be non-trivial
StageCountType,
BuilderScheduleTag,
cute::enable_if_t<cute::is_same_v<KernelMixedTmaCpAsyncWarpSpecialized1SmSm100, BuilderScheduleTag> &&
cute::enable_if_t<(cute::is_same_v<KernelMixedTmaCpAsyncWarpSpecialized1SmSm100, BuilderScheduleTag> ||
cute::is_same_v<KernelMixedTmaCpAsyncWarpSpecialized2SmSm100, BuilderScheduleTag>) &&
(cute::is_same_v<ArchTag, arch::Sm100>
)>
>
@@ -100,7 +101,7 @@ struct CollectiveBuilder<
cute::size<2>(TileShape_MNK{}))));
// Assigning 4 warps for mainloop load of B
static constexpr int NumLoadThreadsCpAsync = 128;
static constexpr int NumLoadThreadsCpAsync = 128 / size(AtomThrID{});
using SmemShapeA_M = decltype(shape_div(shape<0>(TileShape_MNK{}), shape_div(shape<0>(TileShape_MNK{}), size<0>(TileShape_MNK{}) / size(AtomThrID{}))));
@@ -116,7 +117,7 @@ struct CollectiveBuilder<
using GmemCopyAtomB = cute::Copy_Atom<SM80_CP_ASYNC_CACHEGLOBAL_ZFILL<AlignmentTypeB>, ElementB>;
using GmemTiledCopyB = decltype(detail::make_simt_gmem_tiled_copy<
GmemCopyAtomB, NumLoadThreadsCpAsync, AlignmentB, TagToStrideB_t<GmemLayoutBTag>,
decltype(cute::get<1>(TileShape_MNK{})), decltype(cute::get<2>(TileShape_MNK{}))>());
decltype(cute::get<1>(TileShape_MNK{}) / size(AtomThrID{})), decltype(cute::get<2>(TileShape_MNK{}))>());
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{}));
@@ -133,10 +134,19 @@ struct CollectiveBuilder<
static constexpr auto CLCPipelineStorage = sizeof(typename cutlass::PipelineCLCFetchAsync<SchedulerPipelineStageCount, ClusterShape_MNK>::SharedStorage);
// CLC (scheduler) response
static constexpr auto CLCResponseStorage = SchedulerPipelineStageCount * detail::CLCResponseSize;
// Tmem dealloc barrier
static constexpr auto TmemDeallocStorage = sizeof(cutlass::arch::ClusterBarrier);
// MMA trampoline barrier (for 2SM synchronization)
static constexpr auto MmaTrampolineBarrierStorage = sizeof(cutlass::arch::ClusterBarrier);
// Tmem base pointer storage
static constexpr auto TmemBasePtrStorage = sizeof(uint32_t);
// Smem usage that's not part of CollectiveEpilogue::SharedStorage & CollectiveMainloop::SharedStorage
static constexpr auto KernelSmemCarveout = static_cast<int>( AccumulatorPipelineStorage +
CLCPipelineStorage +
CLCResponseStorage);
CLCResponseStorage +
TmemDeallocStorage +
MmaTrampolineBarrierStorage +
TmemBasePtrStorage);
// Reduce SMEM capacity available for buffers considering barrier allocations.
static constexpr int ReducedSmemCapacityBytes = detail::sm100_reduced_smem_capacity_bytes<ArchTag, KernelSmemCarveout>();
using SmemTileShape = cute::Shape<SmemShapeA_M, BlockTileB_N, SmemShapeA_K>;
@@ -188,6 +188,7 @@ struct CollectiveBuilder<
// Dense Gemm / PtrArrayDenseGemm
(
(not cute::is_same_v<KernelMixedTmaCpAsyncWarpSpecialized1SmSm100, BuilderScheduleTag>) &&
(not cute::is_same_v<KernelMixedTmaCpAsyncWarpSpecialized2SmSm100, BuilderScheduleTag>) &&
(not cute::is_same_v<KernelWarpSpecialized1SmSm100, BuilderScheduleTag>) &&
(cute::is_base_of_v<KernelScheduleSm100DenseGemm, BuilderScheduleTag> ||
cute::is_same_v<KernelScheduleAuto, BuilderScheduleTag>)) &&
@@ -503,6 +503,7 @@ check_input_datatypes() {
|| (cute::is_same_v<BuilderScheduleTag, KernelTmaWarpSpecialized1SmBlockScaledSm100>)
|| (cute::is_same_v<BuilderScheduleTag, KernelTmaWarpSpecialized2SmBlockScaledSm100>)
|| (cute::is_same_v<BuilderScheduleTag, KernelMixedTmaCpAsyncWarpSpecialized1SmBlockScaledSm100>)
|| (cute::is_same_v<BuilderScheduleTag, KernelMixedTmaCpAsyncWarpSpecialized2SmBlockScaledSm100>)
// SM100 BS ptr_array
|| (cute::is_same_v<BuilderScheduleTag, KernelSchedulePtrArrayBlockScaledGemmSm100>)
|| (cute::is_same_v<BuilderScheduleTag, KernelPtrArrayTmaWarpSpecialized1SmBlockScaledSm100>)
@@ -103,11 +103,6 @@ struct CollectiveMma<
using TiledMma = TiledMma_;
using AtomThrShapeMNK = Shape<decltype(shape<0>(typename TiledMma::ThrLayoutVMNK{})), _1, _1>;
// Statically asserting to ensure only 1x1x1 cluster shape & 1sm setup is received
static_assert(size(AtomThrShapeMNK{}) == 1, "Lower alignment SM100 GEMM only supports 1SM MMA");
static_assert(size(ClusterShape{}) == 1, "CPASYNC does not support multicast so the cluster shape is restricted to 1, 1, 1");
static_assert(size(typename TiledMma::AtomThrID{}) == 1);
using DispatchPolicy = MainloopSm100UmmaMixedTmaCpAsyncWarpSpecializedBlockScaled<
Stages,
@@ -132,10 +127,13 @@ struct CollectiveMma<
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 LoadShapeA_MK = decltype(select<0,2>(TileShape{}));
using LoadShapeB_NK = decltype(select<1,2>(TileShape{}));
using LoadShapeB_NK = decltype(make_shape(
get<1>(TileShape{}) / size(AtomThrShapeMNK{}), get<2>(TileShape{})
));
// CtaShape_MNK is queried from collective in all kernel layers
using CtaShape_MNK = TileShape;
using CtaShape_MNK = decltype(shape_div(TileShape{}, AtomThrShapeMNK{}));
static_assert(shape<1>(CtaShape_MNK{}) == 192 or shape<1>(CtaShape_MNK{}) == 64 or
shape<1>(CtaShape_MNK{}) == 128 or shape<1>(CtaShape_MNK{}) == 256,
"Cta N should be one of 64/128/192/256");
@@ -320,6 +318,7 @@ struct CollectiveMma<
// Device side kernel params
struct Params {
static_assert(cute::is_static_v<ClusterShape>, "`ClusterShape` must be static in mixed TMA cpasync kernel.");
using ClusterLayout_VMNK = decltype(tiled_divide(make_layout(ClusterShape{}),
make_tile(typename TiledMma::AtomThrID{})));
using ClusterLayoutSfb_VMNK = decltype(tiled_divide(make_layout(ClusterShape{}),
@@ -451,6 +450,25 @@ struct CollectiveMma<
bool implementable = true;
static constexpr bool IsDynamicCluster = !cute::is_static_v<ClusterShape>;
constexpr bool IsBlockscaled = !cute::is_void_v<ElementSF>;
if constexpr (IsBlockscaled) {
if constexpr (IsDynamicCluster) {
implementable &= cutlass::detail::preferred_cluster_can_implement<AtomThrShapeMNK>(args.hw_info.cluster_shape, args.hw_info.cluster_shape_fallback);
// Special cluster shape check for scale factor multicasts. Due to limited size of scale factors, we can't multicast among
// more than 4 CTAs
implementable &= (args.hw_info.cluster_shape.x <= 4 && args.hw_info.cluster_shape.y <= 4 &&
args.hw_info.cluster_shape_fallback.x <= 4 && args.hw_info.cluster_shape_fallback.y <= 4);
}
else {
// Special cluster shape check for scale factor multicasts. Due to limited size of scale factors, we can't multicast among
// more than 4 CTAs
implementable &= ((size<0>(ClusterShape{}) <= 4) && (size<1>(ClusterShape{}) <= 4));
}
}
implementable = implementable && cutlass::detail::check_alignment<min_tma_aligned_elements_A>(cute::make_shape(M,K,L), StrideA{});
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Problem Size doesn't meet the minimum alignment requirements for TMA.\n");
@@ -567,7 +585,7 @@ struct CollectiveMma<
Tensor gSFB_nkl = local_tile(mSFB_nkl, TileShape_SF{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (TILE_N,TILE_K,n,k,l)
ThrMMA cta_mma = TiledMma{}.get_slice(0);
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 sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (MMA,MMA_M,MMA_K,PIPE)
@@ -582,9 +600,10 @@ struct CollectiveMma<
// Define the CTA-in-cluster Layout and Coord
Layout cta_layout_mnk = make_layout(ClusterShape{});
Layout cta_layout_vmnk = tiled_divide(cta_layout_mnk, make_tile(typename TiledMma::AtomThrID{}));
auto cta_coord_vmnk = cta_layout_vmnk.get_flat_coord(0);
uint32_t cta_rank_in_cluster = static_cast<uint32_t>(cute::block_rank_in_cluster());
auto cta_coord_vmnk = cta_layout_vmnk.get_flat_coord(cta_rank_in_cluster);
Layout cta_layout_sfb_vmnk = tiled_divide(cta_layout_mnk, make_tile(typename TiledMma_SF::AtomThrID{}));
auto cta_coord_sfb_vmnk = cta_layout_sfb_vmnk.get_flat_coord(0);
auto cta_coord_sfb_vmnk = cta_layout_sfb_vmnk.get_flat_coord(cta_rank_in_cluster);
// Project the cta_layout for tma_a along the n-modes
auto [tAgA_mkl, tAsA] = tma_partition(*observed_tma_load_a_,
@@ -599,10 +618,15 @@ struct CollectiveMma<
get<1>(cta_coord_sfb_vmnk), make_layout(size<1>(cta_layout_sfb_vmnk)),
group_modes<0,3>(sSFB), group_modes<0,3>(tCgSFB_nkl));
uint16_t mcast_mask_a = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk);
uint16_t mcast_mask_sfa = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk);
uint16_t mcast_mask_sfb = create_tma_multicast_mask<1>(cta_layout_sfb_vmnk, cta_coord_sfb_vmnk);
return cute::make_tuple(
shape<3>(gA_mkl), // for scheduler
tAgA_mkl, tAsA, // for input tensor values
tAgSFA_mkl, tBgSFB_nkl, tAsSFA, tBsSFB // for input scale factor tensor values
tAgSFA_mkl, tBgSFB_nkl, tAsSFA, tBsSFB, // for input scale factor tensor values
mcast_mask_a, mcast_mask_sfa, mcast_mask_sfb
);
}
@@ -628,11 +652,13 @@ struct CollectiveMma<
Tensor mB_nkl = make_tensor(make_gmem_ptr(ptr_B), shape_b, stride_b); //(n,k,l)
// Partition for cpasync
Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k,l)
Tensor tBgB_nkl = flatten(flat_divide(gB_nkl, make_shape(safe_div(size(get<1>(TileShape{})), size(AtomThrShapeMNK{})))));
// Build the coordinate tensors with the same shape as input matrices
Tensor cB_nk = make_identity_tensor(make_shape(N,K));
// Slice the coordinate tensors in the same way as A/B tensor partitioning
Tensor cgB_nk = local_tile(cB_nk, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k)
Tensor ctBgB_nk = flatten(flat_divide(cgB_nk, make_shape(safe_div(size(get<1>(TileShape{})), size(AtomThrShapeMNK{})))));
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), LoadSmemLayoutB{});
@@ -642,7 +668,7 @@ struct CollectiveMma<
auto thr_copy_b = gmem_to_smem_b_tiled_copy.get_slice(thread_idx);
return cute::make_tuple(
gB_nkl, cgB_nk, sB,
tBgB_nkl, ctBgB_nk, sB,
// problem_shape_MNKL,
gmem_to_smem_b_tiled_copy, thr_copy_b);
}
@@ -687,13 +713,13 @@ struct CollectiveMma<
auto tiled_copy_s2t_SFA = make_utccp_copy(UtccpOp{}, tCtSFA_compact);
auto tiled_copy_s2t_SFB = make_utccp_copy(UtccpOp{}, tCtSFB_compact);
auto thr_copy_s2t_SFA = tiled_copy_s2t_SFA.get_slice(0);
auto thr_copy_s2t_SFA = tiled_copy_s2t_SFA.get_slice(blockIdx.x % size(AtomThrID{}));
auto thr_tCsSFA_compact_s2t_ = thr_copy_s2t_SFA.partition_S(tCsSFA_compact);
// SMEM to TMEM copy operation requires source SMEM operand to be an SMEM descriptor
auto thr_tCsSFA_compact_s2t = get_utccp_smem_desc_tensor<UtccpOp>(thr_tCsSFA_compact_s2t_);
auto thr_tCtSFA_compact_s2t = thr_copy_s2t_SFA.partition_D(tCtSFA_compact);
auto thr_copy_s2t_SFB = tiled_copy_s2t_SFB.get_slice(0);
auto thr_copy_s2t_SFB = tiled_copy_s2t_SFB.get_slice(blockIdx.x % size(AtomThrID{}));
auto thr_tCsSFB_compact_s2t_ = thr_copy_s2t_SFB.partition_S(tCsSFB_compact);
// SMEM to TMEM copy operation requires source SMEM operand to be an SMEM descriptor
auto thr_tCsSFB_compact_s2t = get_utccp_smem_desc_tensor<UtccpOp>(thr_tCsSFB_compact_s2t_);
@@ -745,7 +771,9 @@ struct CollectiveMma<
auto [k_tiles,
tAgA_mkl, tAsA,
tAgSFA_mkl, tBgSFB_nkl, tAsSFA, tBsSFB] = load_inputs;
tAgSFA_mkl, tBgSFB_nkl,
tAsSFA, tBsSFB,
mcast_mask_a, mcast_mask_sfa, mcast_mask_sfb] = 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));
@@ -768,9 +796,9 @@ struct CollectiveMma<
barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
if (cute::elect_one_sync()) {
copy(observed_tma_load_a_->with(*tma_barrier), tAgA(_,*k_tile_iter), tAsA(_,write_stage));
copy(observed_tma_load_sfa_->with(*tma_barrier), tAgSFA(_,*k_tile_iter), tAsSFA(_,write_stage));
copy(observed_tma_load_sfb_->with(*tma_barrier), tBgSFB(_,*k_tile_iter), tBsSFB(_,write_stage));
copy(observed_tma_load_a_->with(*tma_barrier, mcast_mask_a), tAgA(_,*k_tile_iter), tAsA(_,write_stage));
copy(observed_tma_load_sfa_->with(*tma_barrier, mcast_mask_sfa), tAgSFA(_,*k_tile_iter), tAsSFA(_,write_stage));
copy(observed_tma_load_sfb_->with(*tma_barrier, mcast_mask_sfb), tBgSFB(_,*k_tile_iter), tBsSFB(_,write_stage));
}
--k_tile_count;
@@ -821,10 +849,12 @@ struct CollectiveMma<
auto [M,N,K,L] = effective_shape;
auto peer_cta_idx = get<0>(cta_coord_mnkl) % size(AtomThrShapeMNK{});
// Slice out the work coord from partitioned tensors
Tensor gB_in = tBgB_nkl(_, _, get<1>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl));
Tensor gB_in = tBgB_nkl(_, peer_cta_idx, _, get<1>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl));
// Repeat slicing out coordinate tensor exactly the same as input tensor does
Tensor cgB_nk_in = cgB_nk(_, _, get<1>(cta_coord_mnkl), _);
Tensor cgB_nk_in = cgB_nk(_, peer_cta_idx, _, get<1>(cta_coord_mnkl), _);
auto k_residue = K - size<1>(gB_in) * size<2>(gB_in); // K - BLK_K * k is negative
@@ -865,7 +895,7 @@ struct CollectiveMma<
copy_if(gmem_to_smem_b_tiled_copy, tBpB, tBgB(_,_,_,*k_tile_iter), tBsB(_,_,_,write_stage));
mainloop_pipeline.producer_commit(mainloop_pipe_producer_state, cutlass::arch::cpasync_barrier_arrive);
mainloop_pipeline.producer_commit_local(mainloop_pipe_producer_state, cutlass::arch::cpasync_barrier_arrive);
--k_tile_count;
++k_tile_iter;
++mainloop_pipe_producer_state;
@@ -892,7 +922,7 @@ struct CollectiveMma<
--k_tile_count;
// UNLOCK mainloop_pipe_producer_state
mainloop_pipeline.producer_commit(mainloop_pipe_producer_state, cutlass::arch::cpasync_barrier_arrive);
mainloop_pipeline.producer_commit_local(mainloop_pipe_producer_state, cutlass::arch::cpasync_barrier_arrive);
// Advance mainloop_pipe_producer_state
++mainloop_pipe_producer_state;
@@ -935,7 +965,11 @@ struct CollectiveMma<
cute::tuple<cute::Tensor<FrgEngine, FrgLayout>> const& accumulators_pair,
cute::tuple<TMmaParams...> const& mma_inputs,
CtaTileCoord cta_tile_coord,
int k_tile_count
int k_tile_count,
bool is_mma_leader_cta,
uint32_t mma_peer_cta_rank,
arch::ClusterBarrier& mma_trampoline_barrier,
uint32_t mma_trampoline_barrier_phase
) {
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)");
@@ -952,6 +986,8 @@ struct CollectiveMma<
auto [mainloop_pipeline_tma, mainloop_pipeline_cpasync, accumulator_pipeline] = pipelines;
auto [mainloop_pipe_tma_consumer_state, mainloop_pipe_cpasync_consumer_state, accumulator_pipe_producer_state] = pipeline_states;
constexpr bool is_2sm = size(AtomThrShapeMNK{}) > 1;
auto tCtSFB_mma = [tCtSFB = tCtSFB, cta_tile_coord]() {
if constexpr (IsCtaN192) {
// If this is an ODD tile, shift the TMEM start address for N=192 case by two words (ignores first 64 columns of SFB)
@@ -973,15 +1009,18 @@ struct CollectiveMma<
}();
// Wait for tmem accumulator buffer to become empty with a flipped phase
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
if (is_mma_leader_cta) {
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
}
//
// PIPELINED MAIN LOOP
//
tiled_mma.accumulate_ = UMMA::ScaleOut::Zero;
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
mainloop_pipeline_tma.consumer_wait(mainloop_pipe_tma_consumer_state);
if (is_mma_leader_cta) {
mainloop_pipeline_tma.consumer_wait(mainloop_pipe_tma_consumer_state);
}
mainloop_pipeline_cpasync.consumer_wait(mainloop_pipe_cpasync_consumer_state);
int read_stage_tma = mainloop_pipe_tma_consumer_state.index();
@@ -992,26 +1031,47 @@ struct CollectiveMma<
copy(tiled_copy_s2t_SFB, thr_tCsSFB_s2t(_,_,_,_,read_stage_tma), thr_tCtSFB_s2t);
}
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.with(tiled_mma.accumulate_,
tCtSFA(_,_,k_block),
tCtSFB_mma(_,_,k_block)),
tCrA(_,_,k_block,read_stage_tma),
tCrB(_,_,k_block,read_stage_cpasync),
accumulators);
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
if (is_mma_leader_cta) {
if constexpr (is_2sm) {
mma_trampoline_barrier.wait(mma_trampoline_barrier_phase);
}
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.with(tiled_mma.accumulate_,
tCtSFA(_,_,k_block),
tCtSFB_mma(_,_,k_block)),
tCrA(_,_,k_block,read_stage_tma),
tCrB(_,_,k_block,read_stage_cpasync),
accumulators);
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
}
} else {
if constexpr (is_2sm) {
mma_trampoline_barrier.arrive(mma_peer_cta_rank);
}
}
mainloop_pipeline_tma.consumer_release(mainloop_pipe_tma_consumer_state);
mainloop_pipeline_cpasync.consumer_release(mainloop_pipe_cpasync_consumer_state);
if constexpr (is_2sm) {
if (is_mma_leader_cta) {
mma_trampoline_barrier.arrive(mma_peer_cta_rank);
} else {
mma_trampoline_barrier.wait(mma_trampoline_barrier_phase);
}
}
if (is_mma_leader_cta) {
mainloop_pipeline_tma.consumer_release(mainloop_pipe_tma_consumer_state);
mainloop_pipeline_cpasync.consumer_release(mainloop_pipe_cpasync_consumer_state);
}
--k_tile_count;
++mainloop_pipe_tma_consumer_state;
++mainloop_pipe_cpasync_consumer_state;
mma_trampoline_barrier_phase ^= 1;
}
return cute::make_tuple(mainloop_pipe_tma_consumer_state, mainloop_pipe_cpasync_consumer_state);
return cute::make_tuple(mainloop_pipe_tma_consumer_state, mainloop_pipe_cpasync_consumer_state, mma_trampoline_barrier_phase);
}
protected:
@@ -56,6 +56,8 @@ 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 <
@@ -63,7 +65,7 @@ template <
int SchedulerPipelineStageCount,
int AccumulatorPipelineStageCount,
class ArchTag_,
class ClusterShape, // Static cluster shape or dynamic (int, int, _1)
class ClusterShape, // Static cluster shape
class TileShape_, // (MmaAtomShapeM, MmaAtomShapeN, TileK)
class ElementA_,
class StrideA_,
@@ -103,12 +105,6 @@ struct CollectiveMma<
using TiledMma = TiledMma_;
using AtomThrShapeMNK = Shape<decltype(shape<0>(typename TiledMma::ThrLayoutVMNK{})), _1, _1>;
// Statically asserting to ensure only 1x1x1 cluster shape & 1sm setup is received
static_assert(size(AtomThrShapeMNK{}) == 1, "Lower alignment SM100 GEMM only supports 1SM MMA");
static_assert(size(ClusterShape{}) == 1, "CPASYNC does not support multicast so the cluster shape is restricted to 1, 1, 1");
static_assert(size(typename TiledMma::AtomThrID{}) == 1);
using DispatchPolicy = MainloopSm100UmmaMixedTmaCpAsyncWarpSpecialized<
Stages,
SchedulerPipelineStageCount,
@@ -124,11 +120,12 @@ struct CollectiveMma<
// Define A and B block shapes
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 LoadShapeA_MK = decltype(select<0,2>(TileShape{}));
using LoadShapeB_NK = decltype(select<1,2>(TileShape{}));
using LoadShapeB_NK = decltype(make_shape(
get<1>(TileShape{}) / size(AtomThrShapeMNK{}), get<2>(TileShape{})
));
// CtaShape_MNK is queried from collective in all kernel layers
using CtaShape_MNK = TileShape;
using CtaShape_MNK = decltype(shape_div(TileShape{}, AtomThrShapeMNK{}));
using ElementA = ElementA_;
using ElementAMma = typename TiledMma::ValTypeA;
@@ -257,6 +254,7 @@ struct CollectiveMma<
// Device side kernel params
struct Params {
static_assert(cute::is_static_v<ClusterShape>, "`ClusterShape` must be static in mixed TMA cpasync kernel.");
using ClusterLayout_VMNK = decltype(tiled_divide(make_layout(ClusterShape{}),
make_tile(typename TiledMma::AtomThrID{})));
@@ -281,8 +279,10 @@ struct CollectiveMma<
CollectiveMma(Params const& params)
: runtime_data_type_a_(params.runtime_data_type_a)
, runtime_data_type_b_(params.runtime_data_type_b) {
observed_tma_load_a_ = &params.tma_load_a;
}
template <class ProblemShape>
@@ -394,9 +394,8 @@ struct CollectiveMma<
/// Set up the data needed by this collective for load.
/// Return tuple element contain
/// gA_mkl - The tiled tensor for input A
/// gB_nkl - The tiled tensor for input B
/// tAsA - partitioned smem tensor for A
/// tBsB - partitioned smem tensor for B
/// mcast_mask_a - tma multicast mask for A
template <class ProblemShape_MNKL>
CUTLASS_DEVICE auto
load_init_tma(
@@ -410,7 +409,7 @@ struct CollectiveMma<
Tensor mA_mkl = observed_tma_load_a_->get_tma_tensor(make_shape(M,K,L));
Tensor gA_mkl = local_tile(mA_mkl, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M, BLK_K, m, k, l)
ThrMMA cta_mma = TiledMma{}.get_slice(0);
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 sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (MMA,MMA_M,MMA_K,PIPE)
@@ -418,16 +417,20 @@ struct CollectiveMma<
// Define the CTA-in-cluster Layout and Coord
Layout cta_layout_mnk = make_layout(ClusterShape{});
Layout cta_layout_vmnk = tiled_divide(cta_layout_mnk, make_tile(typename TiledMma::AtomThrID{}));
auto cta_coord_vmnk = cta_layout_vmnk.get_flat_coord(0);
uint32_t cta_rank_in_cluster = static_cast<uint32_t>(cute::block_rank_in_cluster());
auto cta_coord_vmnk = cta_layout_vmnk.get_flat_coord(cta_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));
uint16_t mcast_mask_a = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk);
return cute::make_tuple(
shape<3>(gA_mkl), // for scheduler
tAgA_mkl, tAsA // for input tensor values
tAgA_mkl, tAsA, // for input tensor values
mcast_mask_a // for TMA multicast
);
}
@@ -451,11 +454,13 @@ struct CollectiveMma<
Tensor mB_nkl = make_tensor(make_gmem_ptr(params.ptr_B), shape_b, stride_b); //(n,k,l)
// Partition for cpasync
Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k,l)
Tensor tBgB_nkl = flatten(flat_divide(gB_nkl, make_shape(safe_div(size(get<1>(TileShape{})), size(AtomThrShapeMNK{})))));
// Build the coordinate tensors with the same shape as input matrices
Tensor cB_nk = make_identity_tensor(make_shape(N,K));
// Slice the coordinate tensors in the same way as A/B tensor partitioning
Tensor cgB_nk = local_tile(cB_nk, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N,BLK_K,n,k)
Tensor ctBgB_nk = flatten(flat_divide(cgB_nk, make_shape(safe_div(size(get<1>(TileShape{})), size(AtomThrShapeMNK{})))));
Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.data()), LoadSmemLayoutB{});
@@ -465,7 +470,7 @@ struct CollectiveMma<
auto thr_copy_b = gmem_to_smem_b_tiled_copy.get_slice(thread_idx);
return cute::make_tuple(
gB_nkl, cgB_nk, sB,
tBgB_nkl, ctBgB_nk, sB,
gmem_to_smem_b_tiled_copy, thr_copy_b);
}
@@ -514,7 +519,8 @@ struct CollectiveMma<
MainloopPipelineTMAState mainloop_pipe_producer_state,
cute::tuple<KTileCount,
GTensorPartitionedA,
STensorA> const& load_inputs,
STensorA,
uint16_t> const& load_inputs,
TileCoordMNKL const& cta_coord_mnkl,
KTileIterator k_tile_iter, int k_tile_count) {
@@ -522,10 +528,11 @@ struct CollectiveMma<
KTileCount k_tiles = get<0>(load_inputs);
GTensorPartitionedA tAgA_mkl = get<1>(load_inputs);
STensorA tAsA = get<2>(load_inputs);
uint16_t mcast_mask_a = get<3>(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));
auto barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
// Issue the Mainloop loads
@@ -542,7 +549,7 @@ struct CollectiveMma<
barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state);
if (cute::elect_one_sync()) {
copy(observed_tma_load_a_->with(*tma_barrier), tAgA(_,*k_tile_iter), tAsA(_,write_stage));
copy(observed_tma_load_a_->with(*tma_barrier, mcast_mask_a), tAgA(_,*k_tile_iter), tAsA(_,write_stage));
}
--k_tile_count;
@@ -583,13 +590,15 @@ struct CollectiveMma<
auto [M,N,K,L] = effective_shape;
auto peer_cta_idx = get<0>(cta_coord_mnkl) % size(AtomThrShapeMNK{});
// Slice out the work coord from partitioned tensors
Tensor gB_in = tBgB_nkl(_, _, get<1>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl));
// Repeat slicing out coordinate tensor exactly the same as input tensor does
Tensor cgB_nk_in = cgB_nk(_, _, get<1>(cta_coord_mnkl), _);
Tensor gB_in = tBgB_nkl(_, peer_cta_idx, _, get<1>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl));
Tensor cgB_nk_in = cgB_nk(_, peer_cta_idx, _, get<1>(cta_coord_mnkl), _);
auto k_residue = K - size<1>(gB_in) * size<2>(gB_in); // K - BLK_K * k is negative
// Repeat slicing out coordinate tensor exactly the same as input tensor does
Tensor gB = gB_in;
Tensor cB = cgB_nk_in;
@@ -627,12 +636,11 @@ struct CollectiveMma<
copy_if(gmem_to_smem_b_tiled_copy, tBpB, tBgB(_,_,_,*k_tile_iter), tBsB(_,_,_,write_stage));
mainloop_pipeline.producer_commit(mainloop_pipe_producer_state, cutlass::arch::cpasync_barrier_arrive);
mainloop_pipeline.producer_commit_local(mainloop_pipe_producer_state, cutlass::arch::cpasync_barrier_arrive);
--k_tile_count;
++k_tile_iter;
++mainloop_pipe_producer_state;
}
// last tile with predication on k to account for residue
// For performance consideration,
// this predicated block for K-tail is only activated when there is k-residue
@@ -654,7 +662,7 @@ struct CollectiveMma<
--k_tile_count;
// UNLOCK mainloop_pipe_producer_state
mainloop_pipeline.producer_commit(mainloop_pipe_producer_state, cutlass::arch::cpasync_barrier_arrive);
mainloop_pipeline.producer_commit_local(mainloop_pipe_producer_state, cutlass::arch::cpasync_barrier_arrive);
// Advance mainloop_pipe_producer_state
++mainloop_pipe_producer_state;
@@ -666,12 +674,6 @@ struct CollectiveMma<
/// Perform a Producer Epilogue to prevent early exit of ctas in a Cluster
CUTLASS_DEVICE void
load_tail_tma(MainloopPipelineTMA mainloop_pipeline, MainloopPipelineTMAState 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);
}
CUTLASS_DEVICE void
@@ -697,7 +699,11 @@ struct CollectiveMma<
cute::tuple<cute::Tensor<FrgEngine, FrgLayout>> const& accumulators_pair,
cute::tuple<TiledMma, FragmentA, FragmentB> const& mma_inputs,
CtaTileCoord cta_tile_coord,
int k_tile_count
int k_tile_count,
bool is_mma_leader_cta,
uint32_t mma_peer_cta_rank,
arch::ClusterBarrier& mma_trampoline_barrier,
uint32_t mma_trampoline_barrier_phase
) {
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)");
@@ -707,37 +713,63 @@ struct CollectiveMma<
auto [mainloop_pipeline_tma, mainloop_pipeline_cpasync, accumulator_pipeline] = pipelines;
auto [mainloop_pipe_tma_consumer_state, mainloop_pipe_cpasync_consumer_state, accumulator_pipe_producer_state] = pipeline_states;
constexpr bool is_2sm = size(AtomThrShapeMNK{}) > 1;
//
// PIPELINED MAIN LOOP
//
tiled_mma.accumulate_ = UMMA::ScaleOut::Zero;
// Wait for tmem accumulator buffer to become empty with a flipped phase
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
if (is_mma_leader_cta) {
accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
}
CUTLASS_PRAGMA_NO_UNROLL
while (k_tile_count > 0) {
mainloop_pipeline_tma.consumer_wait(mainloop_pipe_tma_consumer_state);
if (is_mma_leader_cta) {
mainloop_pipeline_tma.consumer_wait(mainloop_pipe_tma_consumer_state);
}
mainloop_pipeline_cpasync.consumer_wait(mainloop_pipe_cpasync_consumer_state);
int read_stage_tma = mainloop_pipe_tma_consumer_state.index();
int read_stage_cpasync = mainloop_pipe_cpasync_consumer_state.index();
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_tma), tCrB(_,_,k_block,read_stage_cpasync), accumulators);
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
if (is_mma_leader_cta) {
if constexpr (is_2sm) {
mma_trampoline_barrier.wait(mma_trampoline_barrier_phase);
}
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_tma), tCrB(_,_,k_block,read_stage_cpasync), accumulators);
tiled_mma.accumulate_ = UMMA::ScaleOut::One;
}
} else {
if constexpr (is_2sm) {
mma_trampoline_barrier.arrive(mma_peer_cta_rank);
}
}
mainloop_pipeline_tma.consumer_release(mainloop_pipe_tma_consumer_state);
mainloop_pipeline_cpasync.consumer_release(mainloop_pipe_cpasync_consumer_state);
if constexpr (is_2sm) {
if (is_mma_leader_cta) {
mma_trampoline_barrier.arrive(mma_peer_cta_rank);
} else {
mma_trampoline_barrier.wait(mma_trampoline_barrier_phase);
}
}
if (is_mma_leader_cta) {
mainloop_pipeline_tma.consumer_release(mainloop_pipe_tma_consumer_state);
mainloop_pipeline_cpasync.consumer_release(mainloop_pipe_cpasync_consumer_state);
}
--k_tile_count;
++mainloop_pipe_tma_consumer_state;
++mainloop_pipe_cpasync_consumer_state;
mma_trampoline_barrier_phase ^= 1;
}
return cute::make_tuple(mainloop_pipe_tma_consumer_state, mainloop_pipe_cpasync_consumer_state);
return cute::make_tuple(mainloop_pipe_tma_consumer_state, mainloop_pipe_cpasync_consumer_state, mma_trampoline_barrier_phase);
}
protected:
@@ -745,7 +777,6 @@ protected:
typename Params::TMA_A const* observed_tma_load_a_{nullptr};
RuntimeDataTypeA runtime_data_type_a_{};
RuntimeDataTypeB runtime_data_type_b_{};
};
/////////////////////////////////////////////////////////////////////////////////////////////////
+2
View File
@@ -715,6 +715,7 @@ struct KernelTmaWarpSpecialized1SmSm100 final : KernelSchedule1Sm, KernelSchedul
struct KernelTmaWarpSpecialized2SmSm100 final : KernelSchedule2Sm, KernelScheduleSm100DenseGemm {}; // Use for 2SM Dense GEMM Kernels for Collective Mainloop Builder
struct KernelWarpSpecialized1SmSm100 final : KernelSchedule1Sm, KernelScheduleSm100DenseGemm {}; // Use for 1SM Dense GEMM Kernels for Collective Mainloop Builder Without TMA
struct KernelMixedTmaCpAsyncWarpSpecialized1SmSm100 final : KernelSchedule1Sm, KernelScheduleSm100DenseGemm {};
struct KernelMixedTmaCpAsyncWarpSpecialized2SmSm100 final : KernelSchedule2Sm, KernelScheduleSm100DenseGemm {};
///////////////////////////////////////////////////////////////////////////////////////////////////////
// SM100 Ptr-Array Dense GEMM Dispatch Policies
@@ -829,6 +830,7 @@ struct KernelTmaWarpSpecialized2SmMxf4Sm100 final : KernelSchedule2
struct KernelTmaWarpSpecialized1SmMxf8f6f4Sm100 final : KernelSchedule1Sm, KernelScheduleMxf8f6f4Sm100 { };
struct KernelTmaWarpSpecialized2SmMxf8f6f4Sm100 final : KernelSchedule2Sm, KernelScheduleMxf8f6f4Sm100 { };
struct KernelMixedTmaCpAsyncWarpSpecialized1SmBlockScaledSm100 final : KernelSchedule1Sm, KernelScheduleBlockScaledGemmSm100 {};
struct KernelMixedTmaCpAsyncWarpSpecialized2SmBlockScaledSm100 final : KernelSchedule2Sm, KernelScheduleBlockScaledGemmSm100 {};
///////////////////////////////////////////////////////////////////////////////////////////////////////
// SM100 BlockScaled Ptr Array Dense GEMM Dispatch Policies
@@ -416,7 +416,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Mainloop, Epilogue or Scheduler don't meet the requirements for Ptr Array Gemm or Grouped Gemm.\n");
return implementable;
@@ -324,7 +324,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if constexpr (IsDynamicCluster) {
static constexpr int MaxClusterSize = 16;
@@ -357,7 +357,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Mainloop, Epilogue or Scheduler don't meet the requirements for Ptr Array Gemm or Grouped Gemm.\n");
return implementable;
@@ -244,7 +244,8 @@ public:
CUTLASS_TRACE_HOST("to_underlying_arguments(): Setting persistent grid SM count to " << sm_count);
KernelHardwareInfo hw_info{args.hw_info.device_id, sm_count};
KernelHardwareInfo hw_info = args.hw_info;
hw_info.sm_count = sm_count;
// Calculate workspace pointers
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
@@ -286,7 +287,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
static constexpr int MaxClusterSize = 16;
implementable &= size(ClusterShape{}) <= MaxClusterSize;
@@ -59,6 +59,16 @@ namespace cutlass::gemm::kernel {
///////////////////////////////////////////////////////////////////////////////
namespace detail {
template <class T>
struct is_blockscaled_mixed_tma_cpasync : cute::false_type {};
template <int S, int Sched, int Acc, class CS, class AT>
struct is_blockscaled_mixed_tma_cpasync<
MainloopSm100UmmaMixedTmaCpAsyncWarpSpecializedBlockScaled<S, Sched, Acc, CS, AT>
> : cute::true_type {};
} // namespace detail
template <
class ProblemShape_,
class CollectiveMainloop_,
@@ -72,14 +82,18 @@ class GemmUniversal<
TileSchedulerTag_,
cute::enable_if_t<
cutlass::detail::is_kernel_tag_of_v<typename CollectiveMainloop_::DispatchPolicy::Schedule,
KernelMixedTmaCpAsyncWarpSpecializedSm100>>>
KernelMixedTmaCpAsyncWarpSpecializedSm100>
>>
{
public:
using ProblemShape = ProblemShape_;
static constexpr bool IsGroupedGemmKernel = cutlass::gemm::detail::is_moe_problem_shape<ProblemShape>::value;
static constexpr bool IsMoEScheduler = false; // stub for MoE scheduler, which accepts a MoEProblemShape instead of GroupProblemShape
static constexpr bool IsBlockscaled = detail::is_blockscaled_mixed_tma_cpasync<
typename CollectiveMainloop_::DispatchPolicy
>::value;
CUTLASS_HOST_DEVICE
static auto get_problem_shape_gemm(ProblemShape const& shape) {
if constexpr (IsGroupedGemmKernel) {
@@ -156,7 +170,6 @@ public:
using CtaShape_MNK = typename CollectiveMainloop::CtaShape_MNK;
using AtomThrShapeMNK = typename CollectiveMainloop::AtomThrShapeMNK;
static_assert(size(AtomThrShapeMNK{}) == 1, "Lower alignment kernel only supports 1x1x1 cluster shape.");
using TileSchedulerTag = cute::conditional_t<IsGroupedGemmKernel && !IsMoEScheduler, GroupScheduler, TileSchedulerTag_>;
using TileScheduler = typename detail::TileSchedulerSelector<
TileSchedulerTag, ArchTag, CtaShape_MNK, ClusterShape, SchedulerPipelineStageCount, ProblemShape>::Scheduler;
@@ -210,7 +223,8 @@ public:
cutlass::PipelineAsync<SchedulerPipelineStageCount>>;
using CLCPipelineState = typename CLCPipeline::PipelineState;
using TmemAllocator = cute::TMEM::Allocator1Sm;
using TmemAllocator = cute::conditional_t<cute::size(cute::shape<0>(typename TiledMma::ThrLayoutVMNK{})) == 1,
cute::TMEM::Allocator1Sm, cute::TMEM::Allocator2Sm>;
// Kernel level shared memory storage
struct SharedStorage {
@@ -225,6 +239,7 @@ public:
alignas(16) CLCPipelineStorage clc;
alignas(16) AccumulatorPipelineStorage accumulator;
alignas(16) arch::ClusterBarrier tmem_dealloc;
alignas(16) arch::ClusterBarrier mma_trampoline_barrier;
} pipelines;
alignas(16) typename TileScheduler::CLCResponse clc_response[SchedulerPipelineStageCount];
@@ -305,7 +320,8 @@ public:
CUTLASS_TRACE_HOST("to_underlying_arguments(): Setting persistent grid SM count to " << sm_count);
KernelHardwareInfo hw_info{args.hw_info.device_id, sm_count};
KernelHardwareInfo hw_info = args.hw_info;
hw_info.sm_count = sm_count;
// Calculate workspace pointers
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
@@ -371,7 +387,7 @@ public:
auto problem_shape_gemm = get_problem_shape_gemm(args.problem_shape);
implementable &= CollectiveMainloop::can_implement(problem_shape_gemm, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(problem_shape_gemm, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
static constexpr int MaxClusterSize = 16;
implementable &= size(ClusterShape{}) <= MaxClusterSize;
@@ -484,11 +500,13 @@ public:
auto cluster_shape = ClusterShape{};
constexpr int cluster_size = size(ClusterShape{});
int cta_rank_in_cluster = cute::block_rank_in_cluster();
bool is_first_cta_in_cluster = cta_rank_in_cluster == 0;
int cta_coord_v = cta_rank_in_cluster % size<0>(typename TiledMma::AtomThrID{});
bool is_mma_leader_cta = cta_coord_v == 0;
int mma_leader_ctas = size(shape_div(cluster_shape, AtomThrShapeMNK{}));
[[maybe_unused]] uint32_t mma_peer_cta_rank = cta_rank_in_cluster;
constexpr bool has_mma_peer_cta = size(AtomThrShapeMNK{}) == 2;
uint32_t mma_peer_cta_rank = has_mma_peer_cta ? cta_rank_in_cluster ^ 1 : cta_rank_in_cluster;
bool is_mma_leader_cta = cta_coord_v == 0;
[[maybe_unused]] bool is_first_cta_in_cluster = cta_rank_in_cluster == 0;
[[maybe_unused]] uint32_t mma_leader_cta_rank = is_mma_leader_cta? cta_rank_in_cluster : mma_peer_cta_rank;
// Kernel level shared memory storage
SharedStorage& shared_storage = *reinterpret_cast<SharedStorage*>(smem_buf);
@@ -497,12 +515,19 @@ public:
CollectiveMainloop collective_mainloop(params.mainloop);
CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.epilogue);
arch::ClusterBarrier& mma_trampoline_barrier = shared_storage.pipelines.mma_trampoline_barrier;
if (WarpCategory::MMA == warp_category && lane_predicate) {
mma_trampoline_barrier.init(NumMMAThreads);
}
// Do we load source tensor C or other aux inputs
bool is_epi_load_needed = collective_epilogue.is_producer_load_needed();
IsParticipant is_participant = {
(warp_category == WarpCategory::MMA) && is_mma_leader_cta, // mma
(warp_category == WarpCategory::Sched) && is_first_cta_in_cluster, // sched
(warp_category == WarpCategory::MMA), // mma
(warp_category == WarpCategory::Sched)
&& (!IsSchedDynamicPersistent || is_first_cta_in_cluster), // sched
(warp_category == WarpCategory::MainloopLoadTMA), // main_load_tma
(warp_category == WarpCategory::EpilogueLoad) && is_epi_load_needed, // epi_load
(warp_category == WarpCategory::Epilogue), // epilogue
@@ -521,11 +546,27 @@ public:
mainloop_pipeline_tma_params.is_leader = lane_predicate && is_mma_leader_cta && is_participant.main_load_tma;
mainloop_pipeline_tma_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes;
mainloop_pipeline_tma_params.initializing_warp = 0;
MainloopPipelineTMA mainloop_pipeline_tma(shared_storage.pipelines.mainloop.tma,
mainloop_pipeline_tma_params,
cluster_shape,
cute::true_type{}, // Perform barrier init
cute::false_type{}); // Delay mask calculation
MainloopPipelineTMA mainloop_pipeline_tma = [&] () {
if constexpr (IsBlockscaled) {
// If blockscaled, SFB is also multicasted, so we need to wait on the row and column CTAs.
return MainloopPipelineTMA(shared_storage.pipelines.mainloop.tma,
mainloop_pipeline_tma_params,
cluster_shape,
cute::true_type{}, // Perform barrier init
cute::false_type{}); // Delay mask calculation
}
else {
// If not blockscaled, there is no multicast across M mode (i.e. across columsn), so we
// don't need to wait on anything except the row CTAs.
return MainloopPipelineTMA(shared_storage.pipelines.mainloop.tma,
mainloop_pipeline_tma_params,
cluster_shape,
McastDirection::kRow,
cute::true_type{}, // Perform barrier init
cute::false_type{}); // Delay mask calculation
}
}();
// Mainloop Load pipeline (CpAsync)
typename MainloopPipelineCpAsync::Params mainloop_pipeline_cpasync_params;
@@ -606,7 +647,13 @@ public:
accumulator_pipeline_params.producer_arv_count = 1;
accumulator_pipeline_params.consumer_arv_count = size(AtomThrShapeMNK{}) * NumEpilogueThreads;
accumulator_pipeline_params.initializing_warp = 2;
AccumulatorPipeline accumulator_pipeline(shared_storage.pipelines.accumulator, accumulator_pipeline_params, cluster_shape);
AccumulatorPipeline accumulator_pipeline(
shared_storage.pipelines.accumulator,
accumulator_pipeline_params,
cluster_shape,
cute::true_type{}, // Perform barrier init
cute::false_type{} // Delay mask init
);
// Tmem allocator
TmemAllocator tmem_allocator{};
@@ -617,6 +664,11 @@ public:
arch::ClusterBarrier& tmem_deallocation_result_barrier = shared_storage.pipelines.tmem_dealloc;
[[maybe_unused]] uint32_t dealloc_barrier_phase = 0;
if (WarpCategory::MMA == warp_category) {
if (has_mma_peer_cta && lane_predicate) {
tmem_deallocation_result_barrier.init(NumMMAThreads);
}
}
MainloopPipelineTMAState mainloop_pipe_tma_consumer_state;
MainloopPipelineTMAState mainloop_pipe_tma_producer_state = cutlass::make_producer_start_state<MainloopPipelineTMA>();
MainloopPipelineCpAsyncState mainloop_pipe_cpasync_consumer_state;
@@ -639,6 +691,13 @@ public:
pipeline_init_arrive_relaxed(cluster_size);
dim3 block_id_in_cluster = cute::block_id_in_cluster();
if constexpr (IsBlockscaled) {
mainloop_pipeline_tma.init_masks(cluster_shape);
} else {
mainloop_pipeline_tma.init_masks(cluster_shape, McastDirection::kRow);
}
accumulator_pipeline.init_masks(cluster_shape, block_id_in_cluster);
// TileID scheduler
TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, block_id_in_cluster);
typename TileScheduler::WorkTileInfo work_tile_info = scheduler.initial_work_tile_info(cluster_shape);
@@ -713,7 +772,7 @@ public:
auto load_inputs = collective_mainloop.load_init_cpasync(
problem_shape_MNKL, params.mainloop, shared_storage.tensors.mainloop,
scheduler, work_tile_info);
Tensor gA_mkl = get<0>(load_inputs);
Tensor tBgB_nkl = get<0>(load_inputs);
do {
// Get current work tile and fetch next work tile
@@ -722,7 +781,7 @@ public:
auto effective_shape = get_effective_shape(params.problem_shape, work_tile_info);
// Get the number of K tiles to compute for this work as well as the starting K tile offset of the work.
auto k_tile_iter = scheduler.get_k_tile_iterator(work_tile_info, effective_shape, CtaShape_MNK{}, shape<3>(gA_mkl));
auto k_tile_iter = scheduler.get_k_tile_iterator(work_tile_info, effective_shape, CtaShape_MNK{}, shape<4>(tBgB_nkl));
auto k_tile_count = TileScheduler::get_work_k_tile_count(work_tile_info, effective_shape, CtaShape_MNK{});
auto [mainloop_producer_state_next, unused_] = collective_mainloop.load_cpasync(
@@ -756,7 +815,7 @@ public:
}
else if (is_participant.sched) {
if constexpr (IsSchedDynamicPersistent) {
// Whether a new CLC query must be performed.
// See comment below where this variable is updated for a description of
@@ -815,13 +874,13 @@ public:
__syncwarp();
tmem_allocation_result_barrier.arrive();
uint32_t tmem_base_ptr = shared_storage.tmem_base_ptr;
// bulk_tmem.data() = tmem_base_ptr;
collective_mainloop.set_tmem_offsets(tmem_storage, tmem_base_ptr);
uint32_t mma_trampoline_barrier_phase = 0;
// Pass the acc with tuple type since the bgrad kernel change the mma_init API
auto mma_inputs = collective_mainloop.mma_init(params.mainloop,
tmem_storage,
auto mma_inputs = collective_mainloop.mma_init(params.mainloop,
tmem_storage,
shared_storage.tensors.mainloop);
do {
auto effective_shape = get_effective_shape(params.problem_shape, work_tile_info);
@@ -842,8 +901,8 @@ public:
// accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state);
int acc_stage = accumulator_pipe_producer_state.index();
// Tensor accumulators = bulk_tmem(_,_,_,acc_stage);
auto [mainloop_pipe_tma_consumer_state_next_, mainloop_pipe_cpasync_consumer_state_next_] = collective_mainloop.mma(
auto [mainloop_pipe_tma_consumer_state_next_, mainloop_pipe_cpasync_consumer_state_next_, mma_trampoline_barrier_phase_next_] = collective_mainloop.mma(
cute::make_tuple(mainloop_pipeline_tma, mainloop_pipeline_cpasync, accumulator_pipeline),
cute::make_tuple(mainloop_pipe_tma_consumer_state, mainloop_pipe_cpasync_consumer_state, accumulator_pipe_producer_state),
// Pass the acc with tuple type since the bgrad kernel change the mma API
@@ -851,12 +910,20 @@ public:
collective_mainloop.slice_accumulator(tmem_storage, acc_stage),
mma_inputs,
cta_coord_mnkl,
k_tile_count
k_tile_count,
is_mma_leader_cta,
mma_peer_cta_rank,
mma_trampoline_barrier,
mma_trampoline_barrier_phase
);
mainloop_pipe_tma_consumer_state = mainloop_pipe_tma_consumer_state_next_;
mainloop_pipe_cpasync_consumer_state = mainloop_pipe_cpasync_consumer_state_next_;
mma_trampoline_barrier_phase = mma_trampoline_barrier_phase_next_;
accumulator_pipeline.producer_commit(accumulator_pipe_producer_state);
if (is_mma_leader_cta) {
accumulator_pipeline.producer_commit(accumulator_pipe_producer_state);
}
++accumulator_pipe_producer_state;
work_tile_info = next_work_tile_info;
@@ -865,7 +932,15 @@ public:
// Release the right to allocate before deallocations so that the next CTA can rasterize
tmem_allocator.release_allocation_lock();
accumulator_pipeline.producer_tail(accumulator_pipe_producer_state);
if (is_mma_leader_cta) {
accumulator_pipeline.producer_tail(accumulator_pipe_producer_state);
}
if constexpr (has_mma_peer_cta) {
// Leader does wait + arrive, follower does arrive + wait
tmem_deallocation_result_barrier.arrive(mma_peer_cta_rank, not is_mma_leader_cta);
tmem_deallocation_result_barrier.wait(dealloc_barrier_phase);
tmem_deallocation_result_barrier.arrive(mma_peer_cta_rank, is_mma_leader_cta);
}
// Free entire tmem allocation
tmem_allocator.free(tmem_base_ptr, TmemAllocator::Sm100TmemCapacityColumns);
@@ -924,7 +999,6 @@ public:
tmem_allocation_result_barrier.arrive_and_wait();
uint32_t tmem_base_ptr = shared_storage.tmem_base_ptr;
collective_mainloop.set_tmem_offsets(tmem_storage, tmem_base_ptr);
// bulk_tmem.data() = tmem_base_ptr;
bool do_tail_store = false;
do {
@@ -307,7 +307,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if constexpr (IsDynamicCluster) {
static constexpr int MaxClusterSize = 16;
@@ -311,7 +311,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if constexpr (IsDynamicCluster) {
static constexpr int MaxClusterSize = 16;
@@ -310,7 +310,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if constexpr (IsDynamicCluster) {
static constexpr int MaxClusterSize = 16;
@@ -317,7 +317,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if constexpr (IsDynamicCluster) {
static constexpr int MaxClusterSize = 16;
@@ -335,7 +335,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if constexpr (IsDynamicCluster) {
static constexpr int MaxClusterSize = 16;
@@ -35,6 +35,7 @@
#include "cute/int_tuple.hpp"
#include "cutlass/kernel_hardware_info.hpp"
#include "cutlass/arch/config.h"
#include "cutlass/arch/barrier.h"
#include "cutlass/detail/cluster.hpp"
@@ -337,7 +338,7 @@ public:
}
static bool
can_implement(Arguments const& args) {
can_implement(Arguments const& args, KernelHardwareInfo const&) {
return true;
}
@@ -32,6 +32,7 @@
#pragma once
#include "cutlass/kernel_hardware_info.hpp"
#include "cutlass/arch/barrier.h"
#include "cutlass/pipeline/pipeline.hpp"
#include "cutlass/gemm/kernel/sm90_tile_scheduler_group.hpp"
@@ -110,7 +111,7 @@ public:
}
static bool
can_implement(Arguments const& args) {
can_implement(Arguments const& args, KernelHardwareInfo const&) {
return true;
}
@@ -218,8 +218,14 @@ public:
}
static bool
can_implement(Arguments const& args) {
return UnderlyingStreamKScheduler::can_implement(args);
can_implement(Arguments const& args, KernelHardwareInfo const& hw_info) {
if (hw_info.cluster_shape.x != hw_info.cluster_shape_fallback.x ||
hw_info.cluster_shape.y != hw_info.cluster_shape_fallback.y ||
hw_info.cluster_shape.z != hw_info.cluster_shape_fallback.z) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Stream-K scheduler requires cluster shape and fallback cluster shape to be the same.\n");
return false;
}
return UnderlyingStreamKScheduler::can_implement(args, hw_info);
}
CUTLASS_DEVICE
@@ -355,7 +355,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if (!implementable) {
CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Mainloop, Epilogue or Scheduler don't meet the requirements for Ptr Array Gemm or Grouped Gemm.\n");
return implementable;
@@ -314,7 +314,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
if constexpr (IsDynamicCluster) {
implementable &= cutlass::detail::preferred_cluster_can_implement<AtomThrShapeMNK>(args.hw_info.cluster_shape, args.hw_info.cluster_shape_fallback);
@@ -229,7 +229,8 @@ public:
CUTLASS_TRACE_HOST("to_underlying_arguments(): Setting persistent grid SM count to " << sm_count);
KernelHardwareInfo hw_info{args.hw_info.device_id, sm_count};
KernelHardwareInfo hw_info = args.hw_info;
hw_info.sm_count = sm_count;
// Calculate workspace pointers
uint8_t* workspace_ptr = reinterpret_cast<uint8_t*>(workspace);
@@ -275,7 +276,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
+1 -1
View File
@@ -149,7 +149,7 @@ static_assert(is_valid_tile_scheduler, "SM70 kernel does not support specializin
can_implement(Arguments const& args) {
bool mode_implementable = args.mode == GemmUniversalMode::kGemm or
(args.mode == GemmUniversalMode::kBatched && rank(ProblemShape{}) == 4);
return mode_implementable && TileScheduler::can_implement(args.scheduler);
return mode_implementable && TileScheduler::can_implement(args.scheduler, args.hw_info);
}
static size_t
@@ -159,7 +159,7 @@ static_assert(is_valid_tile_scheduler, "SM70 kernel does not support specializin
return implementable;
}
typename ProblemShape::UnderlyingProblemShape problem_shape = args.problem_shape.get_host_problem_shape();
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
@@ -336,7 +336,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
@@ -348,7 +348,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
@@ -167,7 +167,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
@@ -233,7 +233,7 @@ public:
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(transformed_problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
@@ -289,7 +289,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
@@ -287,7 +287,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
@@ -188,7 +188,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
@@ -220,7 +220,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
@@ -233,7 +233,7 @@ public:
}
implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop);
implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue);
implementable &= TileScheduler::can_implement(args.scheduler);
implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info);
return implementable;
}
@@ -241,7 +241,7 @@ public:
}
static bool
can_implement(Arguments const& args) {
can_implement(Arguments const& args, KernelHardwareInfo const&) {
return true;
}
@@ -245,7 +245,7 @@ public:
}
static bool
can_implement(Arguments const& args) {
can_implement(Arguments const& args, KernelHardwareInfo const&) {
// Split count must be positive, and > 1 is only valid for heuristic and split-K decomposition modes
return args.splits >= 1 &&
(args.splits == 1 ||
@@ -126,7 +126,7 @@ public:
CUTLASS_HOST_DEVICE
static bool
can_implement(Arguments const& args) {
can_implement(Arguments const& args, KernelHardwareInfo const&) {
return args.max_swizzle_size >= 0;
}
+5
View File
@@ -91,6 +91,7 @@ struct KernelHardwareInfo {
void const* kernel_ptr,
cudaStream_t stream = nullptr) {
int max_active_clusters = 0;
#if !(defined(__QNX__) && __QNX__ >= 800 && defined(NV_IS_SAFETY))
#if defined(CUTLASS_SM90_CLUSTER_LAUNCH_ENABLED)
ClusterLauncher::LaunchConfig cluster_launch_config = ClusterLauncher::make_cluster_launch_config(
cluster_dims /* minimum grid dim */, cluster_dims, {threads_per_block, 1, 1},
@@ -110,6 +111,10 @@ struct KernelHardwareInfo {
#else
CUTLASS_TRACE_HOST("ClusterLauncher: CUTLASS_SM90_CLUSTER_LAUNCH_ENABLED not defined! Aborting cluster occupancy query.");
return max_active_clusters;
#endif
#else
CUTLASS_TRACE_HOST("ClusterLauncher: cluster launch disabled for QNX 8+ safety builds");
return max_active_clusters;
#endif
}
+13 -1
View File
@@ -550,7 +550,7 @@ public:
using ThreadCategory = typename Impl::ThreadCategory;
using Params = typename Impl::Params;
using McastDirection = McastDirection;
using McastDirection = cutlass::McastDirection;
// Helper function to initialize barriers
static
@@ -820,6 +820,18 @@ public:
impl_.producer_acquire(state, barrier_token);
}
template<class UserDefinedArriveOp>
CUTLASS_DEVICE
void producer_commit_local(PipelineState state, UserDefinedArriveOp&& user_defined_arrive_op) {
cute::forward<UserDefinedArriveOp>(user_defined_arrive_op)(producer_get_barrier(state));
producer_commit_local(state);
}
CUTLASS_DEVICE
void producer_commit_local(PipelineState state) {
impl_.producer_commit(state);
}
template<class UserDefinedArriveOp>
CUTLASS_DEVICE
void producer_commit(PipelineState state, UserDefinedArriveOp&& user_defined_arrive_op) {
+10 -7
View File
@@ -454,18 +454,21 @@ public:
//
// Homebrew read-modify-write
//
Storage original;
Storage updated;
Storage assumed;
#if (__CUDACC_VER_MAJOR__ > 12) || (__CUDACC_VER_MAJOR__ == 12 && __CUDACC_VER_MINOR__ >= 8)
Storage original = __nv_atomic_load_n(ptr_, __NV_ATOMIC_RELAXED);
#else
Storage original = *const_cast<Storage const volatile *>(ptr_);
#endif
do {
original = (*ptr_);
assumed = original;
Storage updated = Storage((assumed & kUpdateMask) | new_bits);
updated = Storage((original & kUpdateMask) | new_bits);
original = atomicCAS(ptr_, assumed, updated);
original = atomicCAS(ptr_, original, updated);
} while (updated != original);
} while (original != assumed);
#else
@@ -204,8 +204,7 @@ public:
void store(Fragment const &frag, TensorCoord const & tile_offset) {
store_with_pointer_offset(
frag,
tile_offset.contiguous() * Shape::kContiguous / ThreadMap::kElementsPerAccess +
tile_offset.strided() * Shape::kStrided * stride_
tile_offset.contiguous() * Shape::kContiguous / ThreadMap::kElementsPerAccess + tile_offset.strided() * Shape::kStrided * stride_
);
}
+2 -2
View File
@@ -35,8 +35,8 @@
#include <string>
#define CUTLASS_MAJOR 4
#define CUTLASS_MINOR 5
#define CUTLASS_PATCH 0
#define CUTLASS_MINOR 4
#define CUTLASS_PATCH 2
#ifdef CUTLASS_VERSIONS_GENERATED
#include "cutlass/version_extended.h"