CUTLASS 3.2.1 (#1113)
* Updates for 3.2.1 release. * Minor fix in gemm op profiler for raster order. * Add scheduler mapping for raster order in the kernels.
This commit is contained in:
@@ -38,9 +38,17 @@
|
||||
#include <cute/atom/copy_traits.hpp>
|
||||
#include <cute/atom/copy_atom.hpp>
|
||||
|
||||
#include <cute/numeric/integral_ratio.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
|
||||
template <class GmemStrides_, class TmaGBasis_, class TmaSwizzle_>
|
||||
struct AuxTmaParams {
|
||||
using GmemStrides = GmemStrides_;
|
||||
GmemStrides g_stride_;
|
||||
};
|
||||
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
///////////////////////////// TMA_LOAD ///////////////////////////////////////
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
@@ -88,14 +96,14 @@ struct Copy_Traits<SM90_TMA_LOAD_OP, NumBitsPerTMA>
|
||||
{
|
||||
static_assert(is_smem<TD>::value, "Expected smem dst for SM90_TMA_LOAD");
|
||||
|
||||
traits.copy_unpack_(raw_pointer_cast(dst.data()), src.data().coord_, tuple_seq<decltype(src.data().coord_)>{});
|
||||
traits.copy_unpack_(cute::raw_pointer_cast(dst.data()), src.data().coord_, tuple_seq<decltype(src.data().coord_)>{});
|
||||
}
|
||||
};
|
||||
|
||||
// The non-executable SM90_TMA_LOAD with tma_desc and no tma_mbar
|
||||
// Use .with(tma_mbar) to construct an executable version
|
||||
template <class NumBitsPerTMA, class GmemStrides>
|
||||
struct Copy_Traits<SM90_TMA_LOAD, NumBitsPerTMA, GmemStrides>
|
||||
template <class NumBitsPerTMA, class AuxParams_>
|
||||
struct Copy_Traits<SM90_TMA_LOAD, NumBitsPerTMA, AuxParams_>
|
||||
{
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
@@ -109,7 +117,8 @@ struct Copy_Traits<SM90_TMA_LOAD, NumBitsPerTMA, GmemStrides>
|
||||
|
||||
// SM90_TMA_LOAD arguments
|
||||
TmaDescriptor tma_desc_;
|
||||
GmemStrides g_stride_;
|
||||
using AuxParams = AuxParams_;
|
||||
AuxParams aux_params_;
|
||||
|
||||
// Return TmaDescriptor/TensorMap
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -133,8 +142,8 @@ struct Copy_Traits<SM90_TMA_LOAD, NumBitsPerTMA, GmemStrides>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_tma_tensor(GShape const& g_shape) const {
|
||||
static_assert(is_congruent<decltype(g_shape), decltype(g_stride_)>::value);
|
||||
return make_counting_tensor(make_layout(g_shape, g_stride_));
|
||||
static_assert(is_congruent<decltype(g_shape), decltype(aux_params_.g_stride_)>::value);
|
||||
return make_counting_tensor(make_layout(g_shape, aux_params_.g_stride_));
|
||||
}
|
||||
|
||||
// Don't try to execute a copy with SM90_TMA_LOAD before calling .with()
|
||||
@@ -190,12 +199,12 @@ struct Copy_Traits<SM90_TMA_LOAD_MULTICAST_OP, NumBitsPerTMA>
|
||||
{
|
||||
static_assert(is_smem<TD>::value, "Expected smem dst for SM90_TMA_LOAD_MULTICAST");
|
||||
|
||||
traits.copy_unpack_(raw_pointer_cast(dst.data()), src.data().coord_, tuple_seq<decltype(src.data().coord_)>{});
|
||||
traits.copy_unpack_(cute::raw_pointer_cast(dst.data()), src.data().coord_, tuple_seq<decltype(src.data().coord_)>{});
|
||||
}
|
||||
};
|
||||
|
||||
template <class NumBitsPerTMA, class GmemStrides>
|
||||
struct Copy_Traits<SM90_TMA_LOAD_MULTICAST, NumBitsPerTMA, GmemStrides>
|
||||
template <class NumBitsPerTMA, class AuxParams_>
|
||||
struct Copy_Traits<SM90_TMA_LOAD_MULTICAST, NumBitsPerTMA, AuxParams_>
|
||||
{
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
@@ -209,7 +218,8 @@ struct Copy_Traits<SM90_TMA_LOAD_MULTICAST, NumBitsPerTMA, GmemStrides>
|
||||
|
||||
// SM90_TMA_LOAD_MULTICAST arguments
|
||||
TmaDescriptor tma_desc_;
|
||||
GmemStrides g_stride_;
|
||||
using AuxParams = AuxParams_;
|
||||
AuxParams aux_params_;
|
||||
|
||||
// Return TmaDescriptor/TensorMap
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -230,8 +240,8 @@ struct Copy_Traits<SM90_TMA_LOAD_MULTICAST, NumBitsPerTMA, GmemStrides>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_tma_tensor(GShape const& g_shape) const {
|
||||
static_assert(is_congruent<decltype(g_shape), decltype(g_stride_)>::value);
|
||||
return make_counting_tensor(make_layout(g_shape, g_stride_));
|
||||
static_assert(is_congruent<decltype(g_shape), decltype(aux_params_.g_stride_)>::value);
|
||||
return make_counting_tensor(make_layout(g_shape, aux_params_.g_stride_));
|
||||
}
|
||||
|
||||
// Don't try to execute a copy with SM90_TMA_LOAD_MULTICAST before calling .with()
|
||||
@@ -248,8 +258,8 @@ struct Copy_Traits<SM90_TMA_LOAD_MULTICAST, NumBitsPerTMA, GmemStrides>
|
||||
//////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
// The executable SM90_TMA_STORE with tma_desc
|
||||
template <class NumBitsPerTMA, class GmemStrides>
|
||||
struct Copy_Traits<SM90_TMA_STORE, NumBitsPerTMA, GmemStrides>
|
||||
template <class NumBitsPerTMA, class AuxParams_>
|
||||
struct Copy_Traits<SM90_TMA_STORE, NumBitsPerTMA, AuxParams_>
|
||||
{
|
||||
using ThrID = Layout<_1>;
|
||||
|
||||
@@ -263,7 +273,8 @@ struct Copy_Traits<SM90_TMA_STORE, NumBitsPerTMA, GmemStrides>
|
||||
|
||||
// SM90_TMA_STORE arguments
|
||||
TmaDescriptor tma_desc_;
|
||||
GmemStrides g_stride_;
|
||||
using AuxParams = AuxParams_;
|
||||
AuxParams aux_params_;
|
||||
|
||||
// Return TmaDescriptor/TensorMap
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -277,8 +288,8 @@ struct Copy_Traits<SM90_TMA_STORE, NumBitsPerTMA, GmemStrides>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
get_tma_tensor(GShape const& g_shape) const {
|
||||
static_assert(is_congruent<decltype(g_shape), decltype(g_stride_)>::value);
|
||||
return make_counting_tensor(make_layout(g_shape, g_stride_));
|
||||
static_assert(is_congruent<decltype(g_shape), decltype(aux_params_.g_stride_)>::value);
|
||||
return make_counting_tensor(make_layout(g_shape, aux_params_.g_stride_));
|
||||
}
|
||||
|
||||
template <class Coord, int... Is>
|
||||
@@ -305,7 +316,7 @@ struct Copy_Traits<SM90_TMA_STORE, NumBitsPerTMA, GmemStrides>
|
||||
static_assert(is_smem<TS>::value, "Expected smem src for SM90_TMA_STORE");
|
||||
//static_assert(is_gmem<TD>::value, "Expected gmem dst for SM90_TMA_STORE"); // TMA spoofed src tensor
|
||||
|
||||
traits.copy_unpack_(raw_pointer_cast(src.data()), dst.data().coord_, tuple_seq<decltype(dst.data().coord_)>{});
|
||||
traits.copy_unpack_(cute::raw_pointer_cast(src.data()), dst.data().coord_, tuple_seq<decltype(dst.data().coord_)>{});
|
||||
}
|
||||
};
|
||||
|
||||
@@ -417,9 +428,78 @@ struct Copy_Traits<SM90_BULK_COPY_AUTO, OpArgs...>
|
||||
|
||||
namespace detail {
|
||||
|
||||
// Use a smem2gmode map to read through the GMEM tensor
|
||||
// and construct a TMA Descriptor for the resulting instruction
|
||||
template <class GEngine, class GLayout,
|
||||
// Custom version of coalesce that greedily combines modes only up to size-256
|
||||
// Look at each element and the back of the stack (in order of priority)
|
||||
// back(NewLayout) get<I>(OldLayout)
|
||||
// s0:d0 _1:d1 => continue
|
||||
// _1:d0 s1:d1 => replace_back s1:d1
|
||||
// s0:d0 s1:s0*d0 => replace_back s0*s1:d0 if s0*s1 <= 256
|
||||
// s0:d0 s1:d1 => append s1:d1
|
||||
//
|
||||
// @pre OldShape and OldStride are flat
|
||||
template <int I, class OldShape, class OldStride, class NewShape, class NewStride>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
coalesce_256_impl(OldShape const& old_shape, OldStride const& old_stride,
|
||||
NewShape const& new_shape, NewStride const& new_stride)
|
||||
{
|
||||
if constexpr (I == rank_v<OldShape>) {
|
||||
// Base case, we're done
|
||||
if constexpr (is_constant<1, NewShape>::value) {
|
||||
return Layout<_1,_0>{};
|
||||
} else {
|
||||
return Layout<NewShape,NewStride>{new_shape,new_stride};
|
||||
}
|
||||
} else if constexpr (is_constant<1, decltype(get<I>(old_shape))>::value) {
|
||||
// shape<I>(layout) == _1, skip it and continue
|
||||
return coalesce_256_impl<I+1>(old_shape, old_stride, new_shape, new_stride);
|
||||
} else if constexpr (is_constant<1, NewShape>::value) {
|
||||
// Replace our shape-1 with anything (Can only happen on input new_shape/new_stride)
|
||||
return coalesce_256_impl<I+1>(old_shape, old_stride, get<I>(old_shape), get<I>(old_stride));
|
||||
} else if constexpr (is_constant<true, decltype(back(new_shape) * back(new_stride) == get<I>(old_stride) &&
|
||||
get<I>(old_shape) * back(new_shape) <= Int<256>{})>::value) {
|
||||
// Merge modes because the shapes and strides match and the merge is 256 or less
|
||||
return coalesce_256_impl<I+1>(old_shape, old_stride,
|
||||
replace_back(new_shape, get<I>(old_shape) * back(new_shape)),
|
||||
new_stride);
|
||||
} else {
|
||||
// Can't replace or merge, so append a new mode
|
||||
return coalesce_256_impl<I+1>(old_shape, old_stride,
|
||||
append(new_shape, get<I>(old_shape)),
|
||||
append(new_stride, get<I>(old_stride)));
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
// Combine all the modes that are possible to combine
|
||||
// Does not respect the profile of the layout, but does preserve total size
|
||||
template <class Shape, class Stride>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
coalesce_256(Layout<Shape,Stride> const& layout)
|
||||
{
|
||||
auto flat_shape = flatten(layout.shape());
|
||||
auto flat_stride = flatten(layout.stride());
|
||||
return coalesce_256_impl<1>(flat_shape, flat_stride, get<0>(flat_shape), get<0>(flat_stride));
|
||||
}
|
||||
|
||||
template <class Engine, class Layout>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
coalesce_256(Tensor<Engine,Layout> const& tensor)
|
||||
{
|
||||
return make_tensor(tensor.data(), coalesce_256(tensor.layout()));
|
||||
}
|
||||
|
||||
|
||||
// Use a smem_inv_h to read through the GMEM tensor
|
||||
// and construct a TMA Descriptor for the resulting instruction
|
||||
// At the same time, construct the Tma Tensor's Stride to generate
|
||||
// the TMA coordinates that the instruction consumes.
|
||||
//
|
||||
template <class TmaInternalType,
|
||||
class GEngine, class GLayout,
|
||||
class SShape, class SStride,
|
||||
int B, int M, int S>
|
||||
CUTE_HOST_RTC
|
||||
@@ -428,63 +508,78 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
Layout<SShape,SStride> const& smem_inv_h, // smem_idx to hier gmode
|
||||
Swizzle<B,M,S> const& swizzle) // Swizzle fn on smem_idx
|
||||
{
|
||||
using T = typename GEngine::value_type;
|
||||
|
||||
// This is the gmem "vector" that corresponds to the smem vector in memory (smem_box_shape):(gmem_prob_stride)
|
||||
Tensor tma_gstride = recast<T>(gtensor.compose(smem_inv_h));
|
||||
|
||||
// If the sizes of smem_inv_h and tma_gstride don't match, then a non-trivial recast was performed.
|
||||
// In that case, require that the recasted modes all have size-1 so TMA can identity them and skip them.
|
||||
for_each(zip(flatten(shape(smem_inv_h)), flatten(shape(tma_gstride))), [] (auto s_and_g) {
|
||||
auto [s,g] = s_and_g;
|
||||
CUTE_STATIC_ASSERT_V(s == g or g == Int<1>{},
|
||||
"A non-trivial recast was performed, but TMA cannot identify which modes to leave out.");
|
||||
});
|
||||
// The smem vector is the same units as gtensor, so compose first and then recast
|
||||
// tma_val_idx:gmem_strides
|
||||
Tensor tile_gstride = recast<TmaInternalType>(gtensor.compose(smem_inv_h));
|
||||
// Coalesce modes up to size-256 (the maximum TMA box extent in units of TmaInternalType)
|
||||
// tma_box_shape:gmem_strides
|
||||
Tensor tma_gstride = coalesce_256(tile_gstride);
|
||||
|
||||
// Perform the tiling to the gmem vector again, but with indirections to the gtensor modes
|
||||
auto gbasis = make_identity_layout(shape(gtensor));
|
||||
auto tma_gbasis_tile_tmp = gbasis.compose(smem_inv_h);
|
||||
// Instead of the recast (gbasis doesn't have type info), replace the shape with the already-recasted shape and coalesce out any size-1 modes
|
||||
auto tma_gbasis_tile = coalesce(make_layout(shape(tma_gstride), stride(tma_gbasis_tile_tmp)));
|
||||
auto tile_gbasis_tmp = gbasis.compose(smem_inv_h);
|
||||
|
||||
// Instead of the recast (gbasis doesn't have type info), replace the shape with the already-recasted shape
|
||||
// tma_box_shape:gmem_mode
|
||||
auto tile_gbasis = make_layout(shape(tile_gstride), stride(tile_gbasis_tmp));
|
||||
|
||||
// Recast the original tensor for shape inspections
|
||||
auto glayout_T = recast<T>(gtensor).layout();
|
||||
auto gtensor_T = recast<TmaInternalType>(gtensor);
|
||||
|
||||
// Find missing bases that don't belong to a size-1 mode of the recast input
|
||||
// Find missing bases that don't appear in tile_gbasis
|
||||
// NOTE This is essentially ArithmeticTuple complement...
|
||||
// NOTE in persuit of implementing an ArithmeticTuple logical_divide for smem_inv_h
|
||||
auto tma_gbasis_full = fold(zip(flatten(shape(glayout_T)), flatten(stride(gbasis))), tma_gbasis_tile,
|
||||
[](auto tma_g, auto s_and_d) {
|
||||
auto [s,d] = s_and_d;
|
||||
auto k = find(stride(tma_g), d); // Find the basis in tma_gstride
|
||||
if constexpr (decltype(k != rank(tma_g) || is_constant<1, decltype(s)>{})::value) {
|
||||
// If d was found or s is static-1, then don't append
|
||||
return tma_g;
|
||||
// NOTE in pursuit of implementing an ArithmeticTuple logical_divide for smem_inv_h
|
||||
auto tile_gbasis_remaining_stride = filter_tuple(flatten(shape (gtensor_T)), flatten(stride(gtensor_T)),
|
||||
flatten(stride(gbasis)),
|
||||
[&](auto s, auto d, auto e)
|
||||
{
|
||||
if constexpr (is_constant<1, decltype(s)>::value || is_constant<0, decltype(d)>::value) {
|
||||
return cute::tuple<>{}; // If size-1 or stride-0, then don't append
|
||||
} else {
|
||||
// Else, append the missing basis
|
||||
return append(tma_g, make_layout(Int<1>{}, d));
|
||||
using E = decltype(e);
|
||||
auto has_e = any_of(stride(tile_gbasis), [] (auto tb) { return tb == E{}; });
|
||||
if constexpr (decltype(has_e)::value) {
|
||||
return cute::tuple<>{}; // If d was found, then don't append
|
||||
} else {
|
||||
return cute::tuple<E>(e); // Else, this is missing so append
|
||||
}
|
||||
}
|
||||
});
|
||||
auto tile_gbasis_remaining_rank = rank(tile_gbasis_remaining_stride);
|
||||
|
||||
// Group the trailing modes to make this max rank-5
|
||||
// "Coalesce" the tile basis into a compatible shape with the tma
|
||||
auto tma_gbasis_tile = tile_gbasis.compose(make_layout(wrap(shape(tma_gstride))));
|
||||
|
||||
// Append the remaining basis modes that contribute to the TMA with size-1
|
||||
auto tma_gbasis_full = make_layout(tuple_cat(wrap( shape(tma_gbasis_tile)), wrap(repeat<tile_gbasis_remaining_rank>(Int<1>{}))),
|
||||
tuple_cat(wrap(stride(tma_gbasis_tile)), wrap(tile_gbasis_remaining_stride)));
|
||||
|
||||
// Group the trailing modes to make this max rank-5 -- TMA rank limitation
|
||||
// tma_box_shape:gmem_mode
|
||||
auto tma_gbasis = group<cute::min(rank(tma_gbasis_full),4),-1>(tma_gbasis_full);
|
||||
|
||||
#if 0
|
||||
print("gtensor : "); print(gtensor); print("\n");
|
||||
print("smem_inv_h : "); print(smem_inv_h); print("\n");
|
||||
print("gtensor : "); print(gtensor); print("\n");
|
||||
print("tile_gstride : "); print(tile_gstride); print("\n");
|
||||
print("tma_gstride : "); print(tma_gstride); print("\n");
|
||||
print("gbasis : "); print(gbasis); print("\n");
|
||||
print("tma_gb_tile : "); print(tma_gbasis_tile ); print("\n");
|
||||
print("tile_gbasis : "); print(tile_gbasis); print("\n");
|
||||
print("tma_gbasis : "); print(tma_gbasis); print("\n");
|
||||
#endif
|
||||
|
||||
//
|
||||
// TMA desc creation
|
||||
//
|
||||
|
||||
constexpr int tma_dim = decltype(rank(tma_gbasis))::value;
|
||||
|
||||
//
|
||||
// TMA gmem desc info
|
||||
//
|
||||
|
||||
void* gmem_address = (void*) raw_pointer_cast(gtensor.data());
|
||||
void* gmem_address = (void*) raw_pointer_cast(gtensor_T.data());
|
||||
auto gmem_layout = gtensor_T.layout();
|
||||
|
||||
cute::array<uint64_t, 5> gmem_prob_shape = {1,1,1,1,1};
|
||||
cute::array<uint64_t, 5> gmem_prob_stride = {0,0,0,0,0};
|
||||
@@ -492,12 +587,12 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
for_each(make_seq<tma_dim>{}, [&](auto i) {
|
||||
for_each(stride<i>(tma_gbasis), [&](auto ej) {
|
||||
// Problem stride
|
||||
uint64_t stride_j = basis_get(ej, stride(glayout_T)) * sizeof(T);
|
||||
uint64_t stride_j = ceil_div(basis_get(ej, stride(gmem_layout)) * sizeof_bits_v<TmaInternalType>, 8);
|
||||
uint64_t old_stride = gmem_prob_stride[i];
|
||||
gmem_prob_stride[i] = gcd(gmem_prob_stride[i], stride_j);
|
||||
|
||||
// Problem shape
|
||||
uint64_t shape_j = basis_get(ej, shape(glayout_T));
|
||||
uint64_t shape_j = basis_get(ej, shape(gmem_layout));
|
||||
if (gmem_prob_stride[i] != 0) {
|
||||
// Recurrence: g_shape = (s_i - 1) * (d_i / gcd_j d_j) + 1
|
||||
gmem_prob_shape[i] = (gmem_prob_shape[i]-1) * (old_stride / gmem_prob_stride[i])
|
||||
@@ -522,8 +617,8 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
assert(gmem_prob_shape[4] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(gmem_prob_shape[4] <= (uint64_t(1) << 32)); // Size must be max 2^32
|
||||
|
||||
// TMA descriptor does not store the zeroth stride and assumes it is sizeof(T) == one element.
|
||||
assert(gmem_prob_stride[0] == sizeof(T) && "Majorness of smem doesn't match majorness of gmem");
|
||||
// TMA descriptor does not store the zeroth stride and assumes it is 1 (TmaInternalType element).
|
||||
assert(gmem_prob_stride[0] == sizeof(TmaInternalType) && "Majorness of smem doesn't match majorness of gmem");
|
||||
|
||||
assert((gmem_prob_stride[1]) < (uint64_t(1) << 40)); // Stride must be max 2^40
|
||||
assert((gmem_prob_stride[1] & 0b1111) == 0); // Stride must be multiple of 16B (128b)
|
||||
@@ -545,14 +640,16 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
smem_box_shape[i] *= size<i>(tma_gbasis);
|
||||
});
|
||||
|
||||
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8 = 256
|
||||
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8 = 256
|
||||
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8 = 256
|
||||
assert(smem_box_shape[0] >= (uint64_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[0] <= (uint64_t(1) << 8)); // Size must be max 2^8 = 256
|
||||
assert(smem_box_shape[0] >= (uint32_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[0] <= (uint32_t(1) << 8)); // Size must be max 2^8 = 256
|
||||
assert(smem_box_shape[1] >= (uint32_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[1] <= (uint32_t(1) << 8)); // Size must be max 2^8 = 256
|
||||
assert(smem_box_shape[2] >= (uint32_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[2] <= (uint32_t(1) << 8)); // Size must be max 2^8 = 256
|
||||
assert(smem_box_shape[3] >= (uint32_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[3] <= (uint32_t(1) << 8)); // Size must be max 2^8 = 256
|
||||
assert(smem_box_shape[4] >= (uint32_t(1))); // Size must be min 1
|
||||
assert(smem_box_shape[4] <= (uint32_t(1) << 8)); // Size must be max 2^8 = 256
|
||||
|
||||
assert(smem_box_stride[0] >= (uint32_t(1))); // Stride must be min 1
|
||||
assert(smem_box_stride[0] <= (uint32_t(8))); // Stride must be max 2^3 = 8
|
||||
@@ -565,88 +662,101 @@ make_tma_copy_desc(Tensor<GEngine,GLayout> const& gtensor, // The original GM
|
||||
assert(smem_box_stride[4] >= (uint32_t(1))); // Stride must be min 1
|
||||
assert(smem_box_stride[4] <= (uint32_t(8))); // Stride must be max 2^3 = 8
|
||||
|
||||
//
|
||||
// Construct the descriptor
|
||||
//
|
||||
|
||||
TmaDescriptor tma_desc = {0};
|
||||
|
||||
//
|
||||
// TMA general info
|
||||
//
|
||||
|
||||
#if (__CUDACC_VER_MAJOR__ >= 12) && !defined(__CUDACC_RTC__)
|
||||
|
||||
CUtensorMapDataType tma_format = TMA::to_CUtensorMapDataType<T>();
|
||||
CUtensorMapInterleave tma_interleave = CU_TENSOR_MAP_INTERLEAVE_NONE;
|
||||
CUtensorMapL2promotion tma_l2Promotion = CU_TENSOR_MAP_L2_PROMOTION_L2_128B;
|
||||
CUtensorMapFloatOOBfill tma_oobFill = CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE;
|
||||
|
||||
// TMA smem swizzle type
|
||||
CUtensorMapSwizzle smem_swizzle = TMA::to_CUtensorMapSwizzle(get_tma_swizzle_bits(swizzle));
|
||||
CUresult result = cuTensorMapEncodeTiled(
|
||||
&tma_desc,
|
||||
tma_format,
|
||||
tma_dim,
|
||||
gmem_address,
|
||||
gmem_prob_shape.data(),
|
||||
gmem_prob_stride.data() + 1, // gmem_prob_stride[0] implicitly 1
|
||||
smem_box_shape.data(),
|
||||
smem_box_stride.data(),
|
||||
tma_interleave,
|
||||
smem_swizzle,
|
||||
tma_l2Promotion,
|
||||
tma_oobFill);
|
||||
|
||||
if (result != CUDA_SUCCESS) {
|
||||
std::cerr << "TMA Desc Addr: " << &tma_desc
|
||||
<< "\nformat " << tma_format
|
||||
<< "\ndim " << tma_dim
|
||||
<< "\ngmem_address " << gmem_address
|
||||
<< "\nglobalDim " << gmem_prob_shape
|
||||
<< "\nglobalStrides " << gmem_prob_stride
|
||||
<< "\nboxDim " << smem_box_shape
|
||||
<< "\nelementStrides " << smem_box_stride
|
||||
<< "\ninterleave " << tma_interleave
|
||||
<< "\nswizzle " << smem_swizzle
|
||||
<< "\nl2Promotion " << tma_l2Promotion
|
||||
<< "\noobFill " << tma_oobFill << std::endl;
|
||||
std::cerr << "Error: Failed to initialize the TMA descriptor " << result << std::endl;
|
||||
assert(false);
|
||||
}
|
||||
|
||||
#endif // (__CUDACC_VER_MAJOR__ >= 12) && !defined(__CUDACC_RTC__)
|
||||
//
|
||||
// Construct the descriptor
|
||||
//
|
||||
|
||||
TmaDescriptor tma_desc = {0};
|
||||
|
||||
//
|
||||
// TMA general info
|
||||
//
|
||||
|
||||
#if (__CUDACC_VER_MAJOR__ >= 12) && !defined(__CUDACC_RTC__)
|
||||
|
||||
CUtensorMapDataType tma_format = TMA::to_CUtensorMapDataType<TmaInternalType>();
|
||||
CUtensorMapInterleave tma_interleave = CU_TENSOR_MAP_INTERLEAVE_NONE;
|
||||
CUtensorMapL2promotion tma_l2Promotion = CU_TENSOR_MAP_L2_PROMOTION_L2_128B;
|
||||
CUtensorMapFloatOOBfill tma_oobFill = CU_TENSOR_MAP_FLOAT_OOB_FILL_NONE;
|
||||
|
||||
// TMA smem swizzle type
|
||||
CUtensorMapSwizzle smem_swizzle = TMA::to_CUtensorMapSwizzle(get_tma_swizzle_bits(swizzle));
|
||||
CUresult result = cuTensorMapEncodeTiled(
|
||||
&tma_desc,
|
||||
tma_format,
|
||||
tma_dim,
|
||||
gmem_address,
|
||||
gmem_prob_shape.data(),
|
||||
gmem_prob_stride.data() + 1, // gmem_prob_stride[0] implicitly 1
|
||||
smem_box_shape.data(),
|
||||
smem_box_stride.data(),
|
||||
tma_interleave,
|
||||
smem_swizzle,
|
||||
tma_l2Promotion,
|
||||
tma_oobFill);
|
||||
|
||||
if (result != CUDA_SUCCESS) {
|
||||
std::cerr << "TMA Desc Addr: " << &tma_desc
|
||||
<< "\nformat " << tma_format
|
||||
<< "\ndim " << tma_dim
|
||||
<< "\ngmem_address " << gmem_address
|
||||
<< "\nglobalDim " << gmem_prob_shape
|
||||
<< "\nglobalStrides " << gmem_prob_stride
|
||||
<< "\nboxDim " << smem_box_shape
|
||||
<< "\nelementStrides " << smem_box_stride
|
||||
<< "\ninterleave " << tma_interleave
|
||||
<< "\nswizzle " << smem_swizzle
|
||||
<< "\nl2Promotion " << tma_l2Promotion
|
||||
<< "\noobFill " << tma_oobFill << std::endl;
|
||||
std::cerr << "Error: Failed to initialize the TMA descriptor " << result << std::endl;
|
||||
assert(false);
|
||||
}
|
||||
|
||||
#endif // (__CUDACC_VER_MAJOR__ >= 12) && !defined(__CUDACC_RTC__)
|
||||
// Finally, get the inverse permutation of the E<i> bases for the mocked gmem stride
|
||||
// NOTE This is essentially ArithmeticTuple inverse...
|
||||
auto gmem_stride_bases = transform_leaf(stride(gbasis), [&](auto ei) {
|
||||
auto si = basis_get(ei, shape(glayout_T));
|
||||
auto di = basis_get(ei, stride(glayout_T));
|
||||
auto tma_gbasis_stride = stride(tma_gbasis);
|
||||
// Find j such that E<i> is in stride<j>(tma_gbasis)
|
||||
[[maybe_unused]] auto j = find_if(tma_gbasis_stride, [&](auto tma_stride_j) { return any_of(tma_stride_j, [&](auto dj) { return dj == ei; }); });
|
||||
// Return the TMA basis this gmode contributes to
|
||||
if constexpr (is_constant<1, decltype(si)>::value || decltype(j == rank(tma_gbasis_stride))::value) {
|
||||
return Int<0>{}; // Return arithmetic identity -- no contribution to the TMA
|
||||
} else
|
||||
if constexpr (decltype(rank<j>(tma_gbasis_stride) == Int<1>{})::value) {
|
||||
return E<j>{}; // We know that the scale factor is Int<1>{}
|
||||
auto si = basis_get(ei, shape(gmem_layout));
|
||||
auto di = basis_get(ei, stride(gmem_layout));
|
||||
if constexpr (is_constant<1, decltype(si)>::value || is_constant<0, decltype(di)>::value) {
|
||||
return Int<0>{}; // If size-1 or stride-0, return arithmetic identity -- no contribution to the TMA
|
||||
} else {
|
||||
return E<j>{} * int32_t(di * sizeof(T) / cute::max(gmem_prob_stride[j], 16));
|
||||
auto tma_gbasis_stride = stride(tma_gbasis);
|
||||
// Find j such that E<i> is in stride<j>(tma_gbasis)
|
||||
using EI = decltype(ei);
|
||||
[[maybe_unused]] auto j = find_if(tma_gbasis_stride, [&](auto tma_stride_j) { return any_of(tma_stride_j, [&](auto dj) { return dj == EI{}; }); });
|
||||
if constexpr (decltype(j == rank(tma_gbasis_stride))::value) {
|
||||
return Int<0>{}; // If not-found, return arithmetic identity -- no contribution to the TMA
|
||||
} else
|
||||
if constexpr (decltype(j == Int<0>{})::value) {
|
||||
auto scale = ratio(size(tma_gstride), size(smem_inv_h)) * basis_get(ei, stride(gtensor));
|
||||
return E<j>{} * scale; // Return TMA Coord basis -- with a recast scale factor
|
||||
} else
|
||||
if constexpr (decltype(rank<j>(tma_gbasis_stride) == Int<1>{})::value) {
|
||||
return E<j>{}; // Return TMA Coord basis -- known scale of Int<1>{}
|
||||
} else {
|
||||
int32_t scale = ceil_div(int32_t(di * sizeof_bits_v<TmaInternalType> / cute::max(gmem_prob_stride[j], 16)), 8);
|
||||
return E<j>{} * scale; // Return TMA Coord basis -- with a dynamic scale factor
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
#if 0
|
||||
print("gmem_stride_bases : "); print(gmem_stride_bases); print("\n");
|
||||
#endif
|
||||
#if 0
|
||||
print("tma_gbasis : "); print(gmem_stride_bases); print("\n");
|
||||
#endif
|
||||
|
||||
return cute::make_tuple(tma_desc, gmem_stride_bases);
|
||||
using AuxParams = AuxTmaParams<decltype(gmem_stride_bases),
|
||||
decltype(tma_gbasis),
|
||||
decltype(swizzle)>;
|
||||
return cute::make_tuple(tma_desc, AuxParams{gmem_stride_bases});
|
||||
}
|
||||
|
||||
// The "logical TMA tid" is a map from the CTA rank to its logical id
|
||||
// within the instruction. It works like a mask or ordering on the
|
||||
// CTAs. For non-multicast TMA, all CTAs should map to 0. For
|
||||
// multicast TMA of size 4, CTAs will be mapped to {0,1,2,3}.
|
||||
template <class CopyOp,
|
||||
template <class TmaInternalType,
|
||||
class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class TShape, class TStride,
|
||||
@@ -657,7 +767,7 @@ make_tma_copy_tiled(CopyOp,
|
||||
Tensor<GEngine,GLayout> const& gtensor, // Full GMEM Tensor
|
||||
SLayout const& slayout, // CTA Tile of SMEM
|
||||
Layout<TShape,TStride> const& cta_t_map, // T: CTA thr idx -> logical TMA tid
|
||||
Layout<VShape,VStride> const& cta_v_map) // V: CTA val idx -> gmem coord
|
||||
Layout<VShape,VStride> const& cta_v_map) // V: CTA val idx -> gmem mode
|
||||
{
|
||||
//
|
||||
// TMA parameter checking
|
||||
@@ -673,18 +783,19 @@ make_tma_copy_tiled(CopyOp,
|
||||
//
|
||||
|
||||
// Invert the smem to get the largest contiguous vector in the smem layout
|
||||
// smem idx -> smem coord
|
||||
auto inv_smem_layout = right_inverse(get_nonswizzle_portion(slayout));
|
||||
// trunc_smem_idx -> trunc_smem_coord
|
||||
|
||||
// Map from smem idx to a gmem mode
|
||||
// Compose with the V-Map to convert smem coord (CTA val idx) to gmem mode
|
||||
// smem idx -> gmem mode
|
||||
auto sidx_to_gmode = coalesce(composition(cta_v_map, inv_smem_layout));
|
||||
|
||||
#if 0
|
||||
print("g_layout : "); print(gtensor.layout()); print("\n");
|
||||
print("g_tensor : "); print(gtensor); print("\n");
|
||||
print("s_layout : "); print(slayout); print("\n");
|
||||
print("cta_t_map : "); print(cta_t_map); print("\n");
|
||||
print("cta_v_map : "); print(cta_v_map); print("\n");
|
||||
print("inv_smem : "); print(inv_smem_layout); print("\n");
|
||||
print("inv_s_layout : "); print(inv_smem_layout); print("\n");
|
||||
print("sidx_to_gmode : "); print(sidx_to_gmode); print("\n");
|
||||
#endif
|
||||
|
||||
@@ -693,9 +804,11 @@ make_tma_copy_tiled(CopyOp,
|
||||
//
|
||||
|
||||
// Generate a TupleBasis for the gtensor
|
||||
// gmem coord -> gmem coord
|
||||
auto glayout_basis = make_identity_layout(shape(gtensor));
|
||||
|
||||
// Tile the modes of gtensor with the truncated cta_v_map o inv_smem_layout_trunc
|
||||
// smem idx -> gmem coord
|
||||
auto tma_layout_full = flatten(composition(glayout_basis, sidx_to_gmode));
|
||||
|
||||
// Truncate any incompatibilities -- no starting in the middle of gmodes
|
||||
@@ -704,61 +817,60 @@ make_tma_copy_tiled(CopyOp,
|
||||
return not is_constant<1,decltype(v)>{};
|
||||
});
|
||||
static_assert(smem_rank > 0, "Could not find a common tile-gmem vectorization. Does the Tile select out major GMEM modes?");
|
||||
// TMA uses a maximum of 5 modes
|
||||
// If the gtensor has more than 5 modes, we need to reserve the last TMA-mode as a "multimode"
|
||||
constexpr int smem_tma_rank = cute::min(int(smem_rank), (rank(tma_layout_full) > Int<5>{} ? 4 : 5));
|
||||
|
||||
// Keep only the static-1 basis modes into gmem
|
||||
auto tma_layout_trunc = take<0,smem_tma_rank>(tma_layout_full);
|
||||
auto tma_layout_trunc = take<0,smem_rank>(tma_layout_full);
|
||||
|
||||
// Split according to the portion each multicast CTA will be responsible for
|
||||
auto tma_layout_vt = logical_divide(tma_layout_trunc, shape_div(size(tma_layout_trunc), cosize(cta_t_map)));
|
||||
// Keep only the portion each multicast CTA will be responsible for
|
||||
auto tma_layout_v = composition(tma_layout_trunc, shape_div(size(tma_layout_trunc), cosize(cta_t_map)));
|
||||
|
||||
#if 0
|
||||
print("glayout_basis : "); print(glayout_basis); print("\n");
|
||||
print("tma_layout_full : "); print(tma_layout_full); print("\n");
|
||||
|
||||
print("tma_layout_trunc: "); print(tma_layout_trunc); print("\n");
|
||||
print("tma_layout_vt : "); print(tma_layout_vt); print("\n");
|
||||
print("tma_layout_v : "); print(tma_layout_v); print("\n");
|
||||
#endif
|
||||
|
||||
//
|
||||
// Construct the TMA Desc and GMEM mode ordering
|
||||
// Construct the TMA Desc and the strides of the TMA Tensor
|
||||
//
|
||||
|
||||
auto [tma_desc, gmem_stride_bases] = detail::make_tma_copy_desc(gtensor, layout<0>(tma_layout_vt), get_swizzle_portion(slayout));
|
||||
auto [tma_desc, aux_params] = detail::make_tma_copy_desc<TmaInternalType>(gtensor,
|
||||
tma_layout_v,
|
||||
get_swizzle_portion(slayout));
|
||||
|
||||
//
|
||||
// Construct the Copy_Traits
|
||||
//
|
||||
|
||||
using T = typename GEngine::value_type;
|
||||
constexpr int num_bits_per_tma = decltype(size<0>(tma_layout_vt))::value * sizeof(T) * 8;
|
||||
using Traits = Copy_Traits<CopyOp, cute::C<num_bits_per_tma>, decltype(gmem_stride_bases)>;
|
||||
constexpr int num_bits_per_tma = decltype(size(tma_layout_trunc))::value * sizeof_bits_v<T>;
|
||||
using Traits = Copy_Traits<CopyOp, cute::C<num_bits_per_tma>, decltype(aux_params)>;
|
||||
using Atom = Copy_Atom<Traits, T>;
|
||||
|
||||
Traits tma_traits{tma_desc, aux_params};
|
||||
|
||||
#if 0
|
||||
print("num_bits : "); print(NumBitsPerTMA{}); print("\n");
|
||||
print("g_stride_bases: "); print(gmem_stride_bases); print("\n");
|
||||
print("num_bits_per_tma : "); print(num_bits_per_tma); print("\n");
|
||||
print("g_stride_bases : "); print(tma_traits.aux_params_.g_stride_); print("\n");
|
||||
#endif
|
||||
|
||||
Traits tma_traits{tma_desc, gmem_stride_bases};
|
||||
|
||||
//
|
||||
// Construct the TiledCopy
|
||||
//
|
||||
|
||||
auto cta_tiler = product_each(shape(cta_v_map));
|
||||
|
||||
// (CTA V, CTA T) -> smem_coord
|
||||
auto layout_vt = composition(inv_smem_layout, make_layout(shape(tma_layout_vt)));
|
||||
// Scale that up to cover all of the smem_coords
|
||||
//
|
||||
// The smem vector might not cover all of the tile,
|
||||
// so multiply it up to cover the entire tile.
|
||||
// "T" here (the parallel index) is a CTA index.
|
||||
auto layout_VT = tile_to_shape(layout_vt, make_shape(size(cta_v_map)/size<1>(layout_vt), size<1>(layout_vt)));
|
||||
// Flip it and change the domain of the T from logical thr to thr_idx
|
||||
auto layout_TV = make_layout(composition(layout<1>(layout_VT), cta_t_map), layout<0>(layout_VT));
|
||||
// CTA V -> smem_coord
|
||||
auto layout_v = composition(inv_smem_layout, size(tma_layout_trunc));
|
||||
auto layout_V = tile_to_shape(make_layout(layout_v), size(cta_v_map));
|
||||
// CTA T -> smem idx
|
||||
auto layout_t = make_layout(cosize(cta_t_map), shape_div(size(tma_layout_trunc), cosize(cta_t_map)));
|
||||
// CTA TID -> smem coord
|
||||
auto layout_T = composition(inv_smem_layout, composition(layout_t, cta_t_map));
|
||||
// Combine with the T mapping
|
||||
auto layout_TV = make_layout(layout_T, layout_V);
|
||||
|
||||
#if 0
|
||||
print("cta_tiler : "); print(cta_tiler); print("\n");
|
||||
@@ -766,8 +878,7 @@ make_tma_copy_tiled(CopyOp,
|
||||
print("layout_TV : "); print(layout_TV); print("\n");
|
||||
#endif
|
||||
|
||||
using T = typename GEngine::value_type;
|
||||
return TiledCopy<Copy_Atom<Traits,T>, decltype(layout_TV), decltype(cta_tiler)>{tma_traits};
|
||||
return TiledCopy<Atom, decltype(layout_TV), decltype(cta_tiler)>{tma_traits};
|
||||
}
|
||||
|
||||
} // end namespace detail
|
||||
@@ -844,6 +955,28 @@ make_tma_copy_tiled(CopyOp,
|
||||
|
||||
copy(tma.with(barrier, mcast_mask), tAgA, tAsA); // copy with supporting TMA params
|
||||
*/
|
||||
template <class TmaInternalType,
|
||||
class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
class CTA_Tile,
|
||||
class Cluster_Size>
|
||||
CUTE_HOST_RTC
|
||||
auto
|
||||
make_tma_copy(CopyOp const& copy_op,
|
||||
Tensor<GEngine,GLayout> const& gtensor,
|
||||
SLayout const& slayout,
|
||||
CTA_Tile const& cta_tile,
|
||||
Cluster_Size const& cluster_size)
|
||||
{
|
||||
return detail::make_tma_copy_tiled<TmaInternalType>(copy_op,
|
||||
gtensor,
|
||||
slayout,
|
||||
make_layout(cluster_size),
|
||||
make_identity_layout(cta_tile));
|
||||
}
|
||||
|
||||
// Explicit defaulting
|
||||
template <class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout,
|
||||
@@ -857,15 +990,14 @@ make_tma_copy(CopyOp const& copy_op,
|
||||
CTA_Tile const& cta_tile,
|
||||
Cluster_Size const& cluster_size)
|
||||
{
|
||||
|
||||
return detail::make_tma_copy_tiled(copy_op,
|
||||
gtensor,
|
||||
slayout,
|
||||
make_layout(cluster_size),
|
||||
make_identity_layout(cta_tile));
|
||||
using TmaInternalType = typename GEngine::value_type;
|
||||
return make_tma_copy<TmaInternalType>(copy_op,
|
||||
gtensor,
|
||||
slayout,
|
||||
cta_tile,
|
||||
cluster_size);
|
||||
}
|
||||
|
||||
// Explicit defaulting
|
||||
template <class CopyOp,
|
||||
class GEngine, class GLayout,
|
||||
class SLayout>
|
||||
|
||||
@@ -155,7 +155,7 @@ struct MMA_Atom<MMA_Traits<Args...>>
|
||||
|
||||
if constexpr (has_dereference<FrgTypeA>::value) {
|
||||
// If the intended FrgTypeA is a view (of the current tensor), forward the whole
|
||||
static_assert(is_same<ValTypeA, typename remove_cvref_t<ATensor>::value_type>::value, "Expecting ValTypeA type");
|
||||
static_assert(is_same<get_raw_type_t<ValTypeA>, typename remove_cvref_t<ATensor>::value_type>::value, "Expecting ValTypeA type");
|
||||
return make_tensor<FrgTypeA>(std::forward<ATensor>(atensor));
|
||||
} else {
|
||||
// Else, the intended FrgTypeA is a value type, construct a new tensor with a fragment layout
|
||||
|
||||
@@ -49,11 +49,11 @@ struct MMA_Traits<SM75_16x8x8_F32F16F16F32_TN>
|
||||
using Shape_MNK = Shape<_16,_8,_8>;
|
||||
using ThrID = Layout<_32>;
|
||||
using ALayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
|
||||
Stride<Stride<_32,_2>,Stride<_16,_1>>>;
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
|
||||
using BLayout = Layout<Shape <Shape < _4,_8>,_2>,
|
||||
Stride<Stride<_16,_1>,_8>>;
|
||||
using CLayout = Layout<Shape <Shape < _4,_8>,Shape < _2,_2>>,
|
||||
Stride<Stride<_32,_2>,Stride<_16,_1>>>;
|
||||
Stride<Stride<_32,_1>,Stride<_16,_8>>>;
|
||||
};
|
||||
|
||||
///////////////////////////////////////////////////////////////////////////////
|
||||
|
||||
Reference in New Issue
Block a user