@@ -39,7 +39,7 @@
|
||||
|
||||
#include <cute/util/type_traits.hpp>
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
#include <cute/tensor_impl.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
@@ -32,7 +32,7 @@
|
||||
|
||||
#include <cute/arch/copy.hpp>
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
#include <cute/tensor_impl.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
@@ -145,4 +145,15 @@ copy_unpack(Copy_Traits<CopyOp,Args...> const& traits,
|
||||
copy_unpack(traits, src, dst);
|
||||
}
|
||||
|
||||
namespace detail {
|
||||
|
||||
template <class CopyOp, class = void>
|
||||
constexpr bool is_prefetch = false;
|
||||
|
||||
template <class CopyOp>
|
||||
constexpr bool is_prefetch<CopyOp, void_t<typename CopyOp::PREFETCH>> = is_same_v<CopyOp, typename CopyOp::PREFETCH>;
|
||||
|
||||
} // end namespace detail
|
||||
|
||||
|
||||
} // end namespace cute
|
||||
|
||||
@@ -39,7 +39,7 @@
|
||||
#include "cute/tensor.hpp"
|
||||
|
||||
#include "cute/algorithm/prefetch.hpp"
|
||||
|
||||
#include "cutlass/fast_math.h"
|
||||
namespace cute
|
||||
{
|
||||
|
||||
@@ -388,18 +388,19 @@ template <class EngineA, class LayoutA,
|
||||
CUTE_HOST
|
||||
auto
|
||||
make_im2col_tma_copy_desc(
|
||||
Tensor<EngineA, LayoutA> const& tensor_cwhdn, // (C,W,H,D,N)
|
||||
uint32_t range_c, // TILE_C
|
||||
uint32_t range_whdn, // TILE_WHDN
|
||||
SmemSwizzle const& smem_swizzle, // Swizzle
|
||||
TMALayout const& tma_layout_vt, // TMA layout
|
||||
LowerCornerStride const& lower_corner_whd, // WHD offset of the "base pointer"
|
||||
UpperCornerStride const& upper_corner_whd, // WHD upper corner
|
||||
LowerPaddingStride const& lower_padding_whd, // WHD lower padding
|
||||
UpperPaddingStride const& upper_padding_whd, // WHD upper padding
|
||||
TraversalStride const& stride_whd, // WHD traversal stride
|
||||
LowerSRTStride const& lower_srt, // SRT offset of the "base pointer"
|
||||
DilationStride const& stride_srt) // SRT stride - dilation
|
||||
Tensor<EngineA, LayoutA> const& tensor_cwhdn, // (C,W,H,D,N)
|
||||
uint32_t range_c, // TILE_C
|
||||
uint32_t range_whdn, // TILE_WHDN
|
||||
SmemSwizzle const& smem_swizzle, // Swizzle
|
||||
TMALayout const& tma_layout_vt, // TMA layout
|
||||
LowerCornerStride const& lower_corner_whd, // WHD offset of the "base pointer"
|
||||
UpperCornerStride const& upper_corner_whd, // WHD upper corner
|
||||
LowerPaddingStride const& lower_padding_whd, // WHD lower padding
|
||||
UpperPaddingStride const& upper_padding_whd, // WHD upper padding
|
||||
TraversalStride const& stride_whd, // WHD traversal stride
|
||||
LowerSRTStride const& lower_srt, // SRT offset of the "base pointer"
|
||||
DilationStride const& stride_srt, // SRT stride - dilation
|
||||
TMA::DescriptorAuxParams const& aux_params = {})
|
||||
{
|
||||
static_assert(is_gmem<EngineA>::value, "Tensor must point to GPU global memory.");
|
||||
using value_type = typename EngineA::value_type;
|
||||
@@ -445,8 +446,8 @@ make_im2col_tma_copy_desc(
|
||||
|
||||
CUtensorMapDataType tma_format = TMA::to_CUtensorMapDataType<value_type>();
|
||||
CUtensorMapInterleave tma_interleave = CU_TENSOR_MAP_INTERLEAVE_NONE;
|
||||
CUtensorMapL2promotion tma_l2Promotion = CU_TENSOR_MAP_L2_PROMOTION_NONE;
|
||||
CUtensorMapFloatOOBfill tma_oob_fill = CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE;
|
||||
CUtensorMapL2promotion tma_l2Promotion = to_CUtensorMapL2promotion(aux_params.l2promo_);
|
||||
CUtensorMapFloatOOBfill tma_oob_fill = to_CUtensorMapFloatOOBfill(aux_params.oobfill_);
|
||||
CUtensorMapSwizzle tma_swizzle = TMA::to_CUtensorMapSwizzle(detail::get_tma_swizzle_bits(smem_swizzle));
|
||||
|
||||
CUresult encode_result = cuTensorMapEncodeIm2col(
|
||||
@@ -498,7 +499,11 @@ make_im2col_tma_copy_desc(
|
||||
|
||||
// For fprop/dgrad kernel, gemm_shapes is ((q, p, z, n), (c, s, r, t))
|
||||
// For wgrad kernel, gemm_shapes is ((c, s, r, t), (q, p, z, n))
|
||||
auto gemm_shapes_common = make_shape(gemm_mn, gemm_k);
|
||||
auto gemm_shapes_common = make_shape(
|
||||
transform_leaf(gemm_mn, [](auto s) {
|
||||
return conditional_return(cute::is_static<decltype(s)>{}, s, cutlass::FastDivmod(s));
|
||||
}),
|
||||
gemm_k);
|
||||
auto gemm_shapes = make_shape(
|
||||
basis_get(stride<0,1>(tma_layout_vt), gemm_shapes_common),
|
||||
basis_get(stride<0,0>(tma_layout_vt), gemm_shapes_common));
|
||||
@@ -554,17 +559,18 @@ template <class CopyOp,
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_atom_im2col(CopyOp,
|
||||
Tensor<GEngine,GLayout> const& gtensor, // Full GMEM Tensor: ((w, h, d, n), c)
|
||||
SLayout const& slayout, // CTA Tile of SMEM, potentially swizzled
|
||||
int32_t const& num_multicast, // The number of CTAs involved in multicasting
|
||||
Layout<VShape,VStride> const& cta_v_map, // V: CTA val idx -> gmem mode
|
||||
LowerCornerStride const& lower_corner_whd,
|
||||
UpperCornerStride const& upper_corner_whd,
|
||||
LowerPaddingStride const& lower_padding_whd,
|
||||
UpperPaddingStride const& upper_padding_whd,
|
||||
TraversalStride const& stride_whd, // traversal stride
|
||||
LowerSRTStride const& lower_srt,
|
||||
DilationStride const& stride_srt) // dilation
|
||||
Tensor<GEngine,GLayout> const& gtensor, // Full GMEM Tensor: ((w, h, d, n), c)
|
||||
SLayout const& slayout, // CTA Tile of SMEM, potentially swizzled
|
||||
int32_t const& num_multicast, // The number of CTAs involved in multicasting
|
||||
Layout<VShape,VStride> const& cta_v_map, // V: CTA val idx -> gmem mode
|
||||
LowerCornerStride const& lower_corner_whd,
|
||||
UpperCornerStride const& upper_corner_whd,
|
||||
LowerPaddingStride const& lower_padding_whd,
|
||||
UpperPaddingStride const& upper_padding_whd,
|
||||
TraversalStride const& stride_whd, // traversal stride
|
||||
LowerSRTStride const& lower_srt,
|
||||
DilationStride const& stride_srt, // dilation
|
||||
TMA::DescriptorAuxParams const& aux_params = {})
|
||||
{
|
||||
//
|
||||
// TMA parameter checking
|
||||
@@ -645,7 +651,8 @@ make_tma_atom_im2col(CopyOp,
|
||||
upper_padding_whd,
|
||||
stride_whd,
|
||||
lower_srt,
|
||||
stride_srt);
|
||||
stride_srt,
|
||||
aux_params);
|
||||
|
||||
//
|
||||
// Construct the Copy_Traits
|
||||
@@ -697,18 +704,19 @@ template <class CopyOp,
|
||||
class DilationStride>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_copy_im2col(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
Layout<TShape,TStride> const& cta_t_map, // CTA tid -> logical TMA tid
|
||||
Layout<VShape,VStride> const& cta_v_map, // CTA vid -> gmem coord
|
||||
LowerCornerStride const& lower_corner_whd,
|
||||
UpperCornerStride const& upper_corner_whd,
|
||||
LowerPaddingStride const& lower_padding_whd,
|
||||
UpperPaddingStride const& upper_padding_whd,
|
||||
TraversalStride const& stride_whd, // traversal stride
|
||||
LowerSRTStride const& lower_srt,
|
||||
DilationStride const& stride_srt) // dilation
|
||||
make_tma_copy_im2col(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
Layout<TShape,TStride> const& cta_t_map, // CTA tid -> logical TMA tid
|
||||
Layout<VShape,VStride> const& cta_v_map, // CTA vid -> gmem coord
|
||||
LowerCornerStride const& lower_corner_whd,
|
||||
UpperCornerStride const& upper_corner_whd,
|
||||
LowerPaddingStride const& lower_padding_whd,
|
||||
UpperPaddingStride const& upper_padding_whd,
|
||||
TraversalStride const& stride_whd, // traversal stride
|
||||
LowerSRTStride const& lower_srt,
|
||||
DilationStride const& stride_srt, // dilation
|
||||
TMA::DescriptorAuxParams const& aux_params = {})
|
||||
{
|
||||
//
|
||||
// TMA parameter checking
|
||||
@@ -719,7 +727,7 @@ make_tma_copy_im2col(CopyOp const& copy_op,
|
||||
|
||||
Copy_Atom atom = make_tma_atom_im2col(copy_op, gtensor, slayout, cosize(cta_t_map), cta_v_map,
|
||||
lower_corner_whd, upper_corner_whd, lower_padding_whd,
|
||||
upper_padding_whd, stride_whd, lower_srt, stride_srt);
|
||||
upper_padding_whd, stride_whd, lower_srt, stride_srt, aux_params);
|
||||
|
||||
//
|
||||
// Construct the TiledCopy
|
||||
|
||||
@@ -124,17 +124,24 @@ struct Copy_Traits<SM90_TMA_LOAD, NumBitsPerTMA, AuxParams_>
|
||||
// Construct an executable SM90_TMA_LOAD with tma_mbar
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
Copy_Traits<SM90_TMA_LOAD_OP, NumBitsPerTMA>
|
||||
with(uint64_t& tma_mbar, [[maybe_unused]] uint16_t const& multicast_mask = 0) const {
|
||||
with(
|
||||
uint64_t& tma_mbar,
|
||||
[[maybe_unused]] uint16_t const& multicast_mask = 0,
|
||||
TMA::CacheHintSm90 const& cache_hint = TMA::CacheHintSm90::EVICT_NORMAL) const {
|
||||
// We accept multicast_mask here to keep the API for both atoms consistent
|
||||
return {{}, {&tma_desc_, &tma_mbar}};
|
||||
return {{}, {&tma_desc_, &tma_mbar, static_cast<uint64_t>(cache_hint)}};
|
||||
}
|
||||
|
||||
// Construct an executable SM90_TMA_LOAD with tma_mbar (temp. overloaded for grouped gemm/ptr array gemm)
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
Copy_Traits<SM90_TMA_LOAD_OP, NumBitsPerTMA>
|
||||
with(TmaDescriptor const* new_tma_desc, uint64_t& tma_mbar, [[maybe_unused]] uint16_t const& multicast_mask = 0) const {
|
||||
with(
|
||||
TmaDescriptor const* new_tma_desc,
|
||||
uint64_t& tma_mbar,
|
||||
[[maybe_unused]] uint16_t const& multicast_mask = 0,
|
||||
TMA::CacheHintSm90 const& cache_hint = TMA::CacheHintSm90::EVICT_NORMAL) const {
|
||||
// We accept multicast_mask here to keep the API for both atoms consistent
|
||||
return {{}, {new_tma_desc, &tma_mbar}};
|
||||
return {{}, {new_tma_desc, &tma_mbar, static_cast<uint64_t>(cache_hint)}};
|
||||
}
|
||||
|
||||
// Generate the TMA coord tensor
|
||||
@@ -171,7 +178,8 @@ struct Copy_Traits<SM90_TMA_LOAD_OP, NumBitsPerTMA>
|
||||
// SM90_TMA_LOAD arguments
|
||||
tuple<
|
||||
TmaDescriptor const*,
|
||||
uint64_t* // smem mbarrier
|
||||
uint64_t*, // smem mbarrier
|
||||
uint64_t // cache hint
|
||||
> const opargs_;
|
||||
};
|
||||
|
||||
@@ -286,6 +294,38 @@ struct Copy_Traits<SM90_TMA_LOAD_MULTICAST_OP, NumBitsPerTMA>
|
||||
///////////////////////////// TMA_STORE //////////////////////////////////////
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// Utility for unpacking TMA_STORE arguments into a CopyOp
|
||||
template <class CopyOp>
|
||||
struct TMA_STORE_Unpack
|
||||
{
|
||||
template <class... Args,
|
||||
class TS, class SLayout,
|
||||
class TD, class DLayout>
|
||||
CUTE_HOST_DEVICE friend constexpr void
|
||||
copy_unpack(Copy_Traits<CopyOp, Args...> const& traits,
|
||||
Tensor<TS,SLayout> const& src,
|
||||
Tensor<TD,DLayout> & dst)
|
||||
{
|
||||
static_assert(is_smem<TS>::value, "Expected smem src for SM90_TMA_STORE");
|
||||
|
||||
void const* const desc_ptr = traits.tma_desc_;
|
||||
void const* const src_ptr = cute::raw_pointer_cast(src.data());
|
||||
auto dst_coord = dst.data().coord_;
|
||||
#if 0
|
||||
auto [c0,c1,c2,c3,c4] = append<5>(dst_coord, 0);
|
||||
printf("THR (%d,%d,%d) BLK (%d,%d,%d) TMACRD (%d,%d,%d,%d,%d) SMEMADDR (%p)\n",
|
||||
threadIdx.x, threadIdx.y, threadIdx.z,
|
||||
blockIdx.x, blockIdx.y, blockIdx.z,
|
||||
int32_t(c0), int32_t(c1), int32_t(c2), int32_t(c3), int32_t(c4), src_ptr);
|
||||
#endif
|
||||
return detail::explode_tuple(detail::CallCOPY<SM90_TMA_STORE>{},
|
||||
make_tuple(desc_ptr, src_ptr), seq<0,1>{},
|
||||
dst_coord, tuple_seq<decltype(dst_coord)>{});
|
||||
}
|
||||
};
|
||||
|
||||
struct SM90_TMA_STORE_OP : SM90_TMA_STORE {};
|
||||
|
||||
// The executable SM90_TMA_STORE with tma_desc
|
||||
template <class NumBitsPerTMA, class AuxParams_>
|
||||
struct Copy_Traits<SM90_TMA_STORE, NumBitsPerTMA, AuxParams_>
|
||||
@@ -343,6 +383,30 @@ struct Copy_Traits<SM90_TMA_STORE, NumBitsPerTMA, AuxParams_>
|
||||
make_tuple(desc_ptr, src_ptr), seq<0,1>{},
|
||||
dst_coord, tuple_seq<decltype(dst_coord)>{});
|
||||
}
|
||||
|
||||
// Construct Copy_Traits executable (w/ swapped out TMA descriptor) for SM90_TMA_STORE (for grouped gemm/ptr array gemm)
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
Copy_Traits<SM90_TMA_STORE_OP, NumBitsPerTMA>
|
||||
with(TmaDescriptor const* new_tma_desc) const {
|
||||
return {{}, new_tma_desc};
|
||||
}
|
||||
};
|
||||
|
||||
// The executable SM90_TMA_STORE with tma_desc
|
||||
template <class NumBitsPerTMA>
|
||||
struct Copy_Traits<SM90_TMA_STORE_OP, NumBitsPerTMA>
|
||||
: TMA_STORE_Unpack<SM90_TMA_STORE_OP>
|
||||
{
|
||||
using ThrID = Layout<_1>;
|
||||
// Map from (src-thr,src-val) to bit
|
||||
using SrcLayout = Layout<Shape<_1,NumBitsPerTMA>>;
|
||||
// Map from (dst-thr,dst-val) to bit
|
||||
using DstLayout = Layout<Shape<_1,NumBitsPerTMA>>;
|
||||
// Reference map from (thr,val) to bit
|
||||
using RefLayout = SrcLayout;
|
||||
|
||||
// SM90_TMA_STORE arguments
|
||||
TmaDescriptor const* tma_desc_;
|
||||
};
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
@@ -1240,14 +1304,14 @@ template <class TmaInternalType = void,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class CTA_Tiler,
|
||||
class Cluster_Size>
|
||||
class Cluster_Size = Int<1>>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_atom(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
CTA_Tiler const& cta_tiler,
|
||||
Cluster_Size const& cluster_size)
|
||||
Cluster_Size const& cluster_size = {})
|
||||
{
|
||||
auto cta_v_tile = make_identity_layout(shape(gtensor)).compose(cta_tiler);
|
||||
// Prefer TmaInternalType if specified. Fallback to GEngine::value_type
|
||||
@@ -1283,8 +1347,8 @@ tma_partition(Copy_Atom<Args...> const& copy_atom,
|
||||
auto layout_V = make_tile(logical_divide(layout_v, tma_layout_v));
|
||||
|
||||
// Append with _ until we cover all Rest... modes
|
||||
auto glayout_V = append<rank_v<decltype(gtensor)>>(layout_V, _);
|
||||
auto slayout_V = append<rank_v<decltype(stensor)>>(layout_V, _);
|
||||
auto glayout_V = append<GLayout::rank>(layout_V, _);
|
||||
auto slayout_V = append<SLayout::rank>(layout_V, _);
|
||||
// Transform tile mode and coalesce
|
||||
Tensor gtensor_v = coalesce(gtensor.compose(glayout_V), Shape<Shape<_1,_1>>{}); // ((TMA,TMA_Iter), Rest...)
|
||||
Tensor stensor_v = coalesce(stensor.compose(slayout_V), Shape<Shape<_1,_1>>{}); // ((TMA,TMA_Iter), Rest...)
|
||||
@@ -1304,8 +1368,8 @@ tma_partition(Copy_Atom<Args...> const& copy_atom,
|
||||
// Offset inside the TMA-mode for the multicast
|
||||
auto multicast_offset = cta_layout(cta_coord) * (size(tma_layout_v) / cosize(cta_layout));
|
||||
auto multicast_coord = make_coord(make_coord(multicast_offset, Int<0>{}));
|
||||
auto scoord = append<SLayout::rank>(multicast_coord, Int<0>{});
|
||||
auto gcoord = append<GLayout::rank>(multicast_coord, Int<0>{});
|
||||
auto scoord = append<SLayout::rank>(multicast_coord, Int<0>{});
|
||||
|
||||
Tensor gresult = domain_offset(gcoord, gtensor_v);
|
||||
Tensor sresult = domain_offset(scoord, stensor_v);
|
||||
@@ -1332,4 +1396,116 @@ create_tma_multicast_mask(CtaLayout const& cta_layout_vmnk,
|
||||
return mcast_mask;
|
||||
}
|
||||
|
||||
////////////////////////////////////
|
||||
// Make TMA copy A/B/C
|
||||
///////////////////////////////////
|
||||
|
||||
template <class TmaInternalType = void,
|
||||
class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class CTA_Tiler,
|
||||
class Cluster_Size>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_copy_A_sm90(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
CTA_Tiler const& cta_tiler,
|
||||
Cluster_Size const& cluster_size)
|
||||
{
|
||||
// Keep only MK modes from MNK
|
||||
auto cta_tiler_mk = remove<1>(cta_tiler);
|
||||
|
||||
// mcast along N mode for this M load, if any
|
||||
auto cluster_size_n = size<1>(cluster_size);
|
||||
|
||||
if constexpr (cute::is_same_v<CopyOp, SM90_TMA_LOAD_IM2COL>) {
|
||||
return make_im2col_tma_copy(copy_op,
|
||||
gtensor,
|
||||
slayout,
|
||||
cta_tiler_mk,
|
||||
cluster_size_n);
|
||||
} else {
|
||||
auto cta_v_tile = make_identity_layout(shape(gtensor)).compose(cta_tiler_mk);
|
||||
auto cta_t_tile = make_layout(cluster_size_n);
|
||||
|
||||
// Prefer TmaInternalType if specified. Fallback to GEngine::value_type
|
||||
using TmaType = conditional_t<is_same<void, TmaInternalType>::value, typename GEngine::value_type, TmaInternalType>;
|
||||
auto tma_copy = detail::make_tma_copy_tiled<TmaType>(copy_op, gtensor, slayout, cta_t_tile, cta_v_tile);
|
||||
return tma_copy;
|
||||
}
|
||||
}
|
||||
|
||||
template <class TmaInternalType = void,
|
||||
class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class CTA_Tiler,
|
||||
class Cluster_Size>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_copy_B_sm90(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
CTA_Tiler const& cta_tiler,
|
||||
Cluster_Size const& cluster_size)
|
||||
{
|
||||
// Keep only NK modes from MNK
|
||||
auto cta_tiler_nk = remove<0>(cta_tiler);
|
||||
|
||||
// mcast along M mode for this N load, if any
|
||||
auto cluster_size_m = size<0>(cluster_size);
|
||||
|
||||
if constexpr (cute::is_same_v<CopyOp, SM90_TMA_LOAD_IM2COL>) {
|
||||
return make_im2col_tma_copy(copy_op,
|
||||
gtensor,
|
||||
slayout,
|
||||
cta_tiler_nk,
|
||||
cluster_size_m);
|
||||
} else {
|
||||
auto cta_v_tile = make_identity_layout(shape(gtensor)).compose(cta_tiler_nk);
|
||||
auto cta_t_tile = make_layout(cluster_size_m);
|
||||
|
||||
// Prefer TmaInternalType if specified. Fallback to GEngine::value_type
|
||||
using TmaType = conditional_t<is_same<void, TmaInternalType>::value, typename GEngine::value_type, TmaInternalType>;
|
||||
auto tma_copy = detail::make_tma_copy_tiled<TmaType>(copy_op, gtensor, slayout, cta_t_tile, cta_v_tile);
|
||||
return tma_copy;
|
||||
}
|
||||
}
|
||||
|
||||
template <class TmaInternalType = void,
|
||||
class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class CTA_Tiler>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_copy_C_sm90(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
CTA_Tiler const& cta_tiler)
|
||||
{
|
||||
// Keep only MN modes from MNK
|
||||
auto cta_tiler_mn = remove<2>(cta_tiler);
|
||||
|
||||
if constexpr (cute::is_same_v<CopyOp, SM90_TMA_LOAD_IM2COL> ||
|
||||
cute::is_same_v<CopyOp, SM90_TMA_STORE_IM2COL>) {
|
||||
return make_im2col_tma_copy(copy_op,
|
||||
gtensor,
|
||||
slayout,
|
||||
cta_tiler_mn,
|
||||
_1{});
|
||||
} else {
|
||||
auto cta_v_tile = make_identity_layout(shape(gtensor)).compose(cta_tiler_mn);
|
||||
|
||||
// No multicast, so only 1 CTA involved
|
||||
auto cta_t_map = Layout<_1,_0>{};
|
||||
|
||||
// Prefer TmaInternalType if specified. Fallback to GEngine::value_type
|
||||
using TmaType = conditional_t<is_same<void, TmaInternalType>::value, typename GEngine::value_type, TmaInternalType>;
|
||||
auto tma_copy = detail::make_tma_copy_tiled<TmaType>(copy_op, gtensor, slayout, cta_t_map, cta_v_tile);
|
||||
return tma_copy;
|
||||
}
|
||||
}
|
||||
} // end namespace cute
|
||||
|
||||
@@ -31,11 +31,9 @@
|
||||
#pragma once
|
||||
|
||||
#include <cute/config.hpp>
|
||||
|
||||
#include <cute/arch/mma.hpp>
|
||||
|
||||
#include <cute/atom/mma_traits.hpp>
|
||||
#include <cute/tensor.hpp>
|
||||
#include <cute/tensor_impl.hpp>
|
||||
#include <cute/util/type_traits.hpp>
|
||||
|
||||
namespace cute {
|
||||
@@ -102,7 +100,7 @@ struct MMA_Atom<MMA_Traits<Args...>>
|
||||
static_assert(BLayout::rank == 1, "Expected rank-1 B tensor");
|
||||
static_assert(CLayout::rank == 1, "Expected rank-1 C tensor");
|
||||
|
||||
return mma_unpack(*this, D, A, B, C);
|
||||
return mma_unpack(static_cast<Traits const&>(*this), D, A, B, C);
|
||||
}
|
||||
|
||||
// Three arguments reproduces C
|
||||
@@ -245,12 +243,9 @@ struct TiledMMA : MMA_Atom
|
||||
thrfrg_C(CTensor&& ctensor) const
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(ctensor) >= Int<2>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<0>(ctensor) % size<0>(TiledShape_MNK{}) == Int<0>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<1>(ctensor) % size<1>(TiledShape_MNK{}) == Int<0>{});
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<0>(PermutationMNK{}),
|
||||
get<1>(PermutationMNK{}));
|
||||
auto t_tile = make_tile(permutation_mnk<0>(),
|
||||
permutation_mnk<1>());
|
||||
auto t_tensor = logical_divide(ctensor, t_tile); // (PermM,PermN)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
@@ -287,12 +282,9 @@ struct TiledMMA : MMA_Atom
|
||||
thrfrg_A(ATensor&& atensor) const
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(atensor) >= Int<2>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<0>(atensor) % size<0>(TiledShape_MNK{}) == Int<0>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<1>(atensor) % size<2>(TiledShape_MNK{}) == Int<0>{});
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<0>(PermutationMNK{}),
|
||||
get<2>(PermutationMNK{}));
|
||||
auto t_tile = make_tile(permutation_mnk<0>(),
|
||||
permutation_mnk<2>());
|
||||
auto t_tensor = logical_divide(atensor, t_tile); // (PermM,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
@@ -329,12 +321,9 @@ struct TiledMMA : MMA_Atom
|
||||
thrfrg_B(BTensor&& btensor) const
|
||||
{
|
||||
CUTE_STATIC_ASSERT_V(rank(btensor) >= Int<2>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<0>(btensor) % size<1>(TiledShape_MNK{}) == Int<0>{});
|
||||
//CUTE_STATIC_ASSERT_V(size<1>(btensor) % size<2>(TiledShape_MNK{}) == Int<0>{});
|
||||
|
||||
// Reorder the tensor for the TiledAtom
|
||||
auto t_tile = make_tile(get<1>(PermutationMNK{}),
|
||||
get<2>(PermutationMNK{}));
|
||||
auto t_tile = make_tile(permutation_mnk<1>(),
|
||||
permutation_mnk<2>());
|
||||
auto t_tensor = logical_divide(btensor, t_tile); // (PermN,PermK)
|
||||
|
||||
// Tile the tensor for the Atom
|
||||
@@ -377,21 +366,23 @@ struct TiledMMA : MMA_Atom
|
||||
// Utility for printing and visualization
|
||||
//
|
||||
|
||||
// The permutation applied to the MNK-mode data
|
||||
template <int I>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
permutation_mnk() const {
|
||||
static_assert(0 <= I && I < 3);
|
||||
auto perm = get<I>(PermutationMNK{});
|
||||
return conditional_return(is_underscore<decltype(perm)>{}, size<I>(AtomShape_MNK{}) * size<I+1>(get_thr_layout_vmnk()), perm);
|
||||
}
|
||||
|
||||
// The size of the MNK-mode
|
||||
template <int I>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
tile_size_mnk() const {
|
||||
static_assert(0 <= I && I < 3);
|
||||
auto core_size = size<I>(AtomShape_MNK{}) * size<I+1>(get_thr_layout_vmnk());
|
||||
[[maybe_unused]] auto perm_size = size<I>(PermutationMNK{});
|
||||
if constexpr (is_underscore<decltype(perm_size)>::value) {
|
||||
return core_size;
|
||||
} else {
|
||||
return cute::max(core_size, perm_size);
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
return size(permutation_mnk<I>());
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
|
||||
@@ -32,7 +32,7 @@
|
||||
|
||||
#include <cute/arch/mma.hpp>
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
#include <cute/tensor_impl.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
@@ -332,9 +332,6 @@ struct DescriptorIterator
|
||||
{
|
||||
return { GmmaDescriptor{desc_ + uint64_t(offset)} };
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE friend void
|
||||
print(DescriptorIterator) { printf("GMMA::DescriptorIterator"); }
|
||||
};
|
||||
|
||||
template <class T>
|
||||
@@ -353,6 +350,11 @@ recast_ptr(DescriptorIterator const& iter) {
|
||||
return iter; // Do nothing, it will still dereference to GmmaDescriptor and decay to uint64_t
|
||||
}
|
||||
|
||||
CUTE_HOST_DEVICE void
|
||||
print(DescriptorIterator) {
|
||||
printf("GMMA::DescriptorIterator");
|
||||
}
|
||||
|
||||
// The GMMA Traits below have custom fragment type flags for their smem desc tensors.
|
||||
// These flags specialize a MakeTensor customization point to correctly make the fragment that is desired.
|
||||
template <GMMA::Major>
|
||||
|
||||
Reference in New Issue
Block a user