More updates for 3.1 (#958)
* Updates for 3.1 * Minor change * doc link fix * Minor updates
This commit is contained in:
@@ -499,7 +499,7 @@ flatten(T const& t)
|
||||
|
||||
namespace detail {
|
||||
|
||||
// Shortcut around tuple_cat for common insert/remove/repeat cases
|
||||
// Shortcut around cute::tuple_cat for common insert/remove/repeat cases
|
||||
template <class T, class X, int... I, int... J, int... K>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
|
||||
@@ -623,7 +623,7 @@ partition_shape_C(TiledMMA<Args...> const& mma, Shape_MN const& shape_MN)
|
||||
auto V = shape<1>(typename TiledMMA<Args...>::AtomLayoutC_TV{});
|
||||
auto M = shape_div(size<0>(shape_MN), size<0>(atomMNK) * size<1>(thrVMNK));
|
||||
auto N = shape_div(size<1>(shape_MN), size<1>(atomMNK) * size<2>(thrVMNK));
|
||||
return tuple_cat(make_shape(V,M,N), take<2,R>(shape_MN));
|
||||
return cute::tuple_cat(make_shape(V,M,N), take<2,R>(shape_MN));
|
||||
}
|
||||
|
||||
template <class... Args, class Shape_MN>
|
||||
@@ -651,7 +651,7 @@ partition_shape_A(TiledMMA<Args...> const& mma, Shape_MK const& shape_MK)
|
||||
auto V = shape<1>(typename TiledMMA<Args...>::AtomLayoutA_TV{});
|
||||
auto M = shape_div(size<0>(shape_MK), size<0>(atomMNK) * size<1>(thrVMNK));
|
||||
auto K = shape_div(size<1>(shape_MK), size<2>(atomMNK) * size<3>(thrVMNK));
|
||||
return tuple_cat(make_shape(V,M,K), take<2,R>(shape_MK));
|
||||
return cute::tuple_cat(make_shape(V,M,K), take<2,R>(shape_MK));
|
||||
}
|
||||
|
||||
template <class... Args, class Shape_NK>
|
||||
@@ -666,7 +666,7 @@ partition_shape_B(TiledMMA<Args...> const& mma, Shape_NK const& shape_NK)
|
||||
auto V = shape<1>(typename TiledMMA<Args...>::AtomLayoutB_TV{});
|
||||
auto N = shape_div(size<0>(shape_NK), size<1>(atomMNK) * size<2>(thrVMNK));
|
||||
auto K = shape_div(size<1>(shape_NK), size<2>(atomMNK) * size<3>(thrVMNK));
|
||||
return tuple_cat(make_shape(V,N,K), take<2,R>(shape_NK));
|
||||
return cute::tuple_cat(make_shape(V,N,K), take<2,R>(shape_NK));
|
||||
}
|
||||
|
||||
//
|
||||
|
||||
@@ -46,8 +46,14 @@ namespace cute
|
||||
|
||||
using dim3 = ::dim3;
|
||||
|
||||
// MSVC doesn't define its C++ version macro to match
|
||||
// its C++ language version. This means that when
|
||||
// building with MSVC, dim3 isn't constexpr-friendly.
|
||||
template <size_t I>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
CUTE_HOST_DEVICE
|
||||
#if ! defined(_MSC_VER)
|
||||
constexpr
|
||||
#endif
|
||||
uint32_t& get(dim3& a)
|
||||
{
|
||||
static_assert(I < 3, "Index out of range");
|
||||
@@ -63,7 +69,10 @@ uint32_t& get(dim3& a)
|
||||
}
|
||||
|
||||
template <size_t I>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
CUTE_HOST_DEVICE
|
||||
#if ! defined(_MSC_VER)
|
||||
constexpr
|
||||
#endif
|
||||
uint32_t const& get(dim3 const& a)
|
||||
{
|
||||
static_assert(I < 3, "Index out of range");
|
||||
@@ -79,7 +88,10 @@ uint32_t const& get(dim3 const& a)
|
||||
}
|
||||
|
||||
template <size_t I>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
CUTE_HOST_DEVICE
|
||||
#if ! defined(_MSC_VER)
|
||||
constexpr
|
||||
#endif
|
||||
uint32_t&& get(dim3&& a)
|
||||
{
|
||||
static_assert(I < 3, "Index out of range");
|
||||
|
||||
Reference in New Issue
Block a user