cutlass 3.9 update (#2255)
* cutlass 3.9 update * rebase * fixes out of shared memory for blockwise Blackwell * doc format * fix issue 2253 * disable host ref by default * fix sm120 smem capacity --------- Co-authored-by: yuzhai <yuzhai@nvidia.com> Co-authored-by: Haicheng Wu <haichengw@nvidia.com>
This commit is contained in:
co-authored by
yuzhai
Haicheng Wu
parent
8e345c5c5b
commit
331a1f5b3f
@@ -51,8 +51,8 @@
|
||||
// but do _not_ include references like int& or float&.
|
||||
// (See std::tie for an example of a tuple of references.)
|
||||
//
|
||||
// Standard-layout types preserve ABI across host-device boundaries.
|
||||
// They are safe to use as device kernel parameters.
|
||||
// Standard-layout types preserve ABI across host-device boundaries. They are safe to use as device kernel parameters.
|
||||
// The standard-layout requirement prevents a more common EBO-based implemented of cute::tuple.
|
||||
//
|
||||
// The cute::tuple is also simplified over the implementations in std::, cuda::std::, and thrust:: by ignoring much of
|
||||
// the conversion SFINAE, special overloading, and avoiding cvref template types.
|
||||
@@ -62,12 +62,15 @@
|
||||
namespace cute
|
||||
{
|
||||
|
||||
namespace detail
|
||||
template <class... T>
|
||||
struct tuple;
|
||||
|
||||
namespace eso
|
||||
{
|
||||
|
||||
// ESO stands for "empty structure optimization."
|
||||
// We use this technique to ensure that cute::tuple
|
||||
// doesn't waste space storing template arguments that have no data (like integral_constant).
|
||||
// We use this technique to ensure that cute::tuple doesn't waste space
|
||||
// storing template arguments that have no data (like integral_constant).
|
||||
// Empty types in the template argument list are not even constructed,
|
||||
// and do not have unique element addresses. Calling `get`
|
||||
// constructs and returns an instance of an empty type on demand.
|
||||
@@ -131,94 +134,92 @@ struct ESO<false, false, First, Rest...> {
|
||||
};
|
||||
|
||||
// Get Nth value from ESO
|
||||
template <size_t N, bool F, bool R, class T, class... Rest>
|
||||
template <class R, size_t N, class S>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
cute::enable_if_t<cute::is_empty<cute::tuple_element_t<N, cute::type_list<T, Rest...>>>::value,
|
||||
cute::tuple_element_t<N, cute::type_list<T, Rest...>>>
|
||||
getv(ESO<F, R, T, Rest...> const&)
|
||||
{
|
||||
return {};
|
||||
}
|
||||
|
||||
template <size_t N, bool F, bool R, class T, class... Rest>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
cute::enable_if_t<not cute::is_empty<cute::tuple_element_t<N, cute::type_list<T, Rest...>>>::value,
|
||||
cute::tuple_element_t<N, cute::type_list<T, Rest...>> const&>
|
||||
getv(ESO<F, R, T, Rest...> const& s)
|
||||
R
|
||||
getr(S&& s) noexcept
|
||||
{
|
||||
if constexpr (N == 0) {
|
||||
return static_cast<T const&>(s.first_);
|
||||
return static_cast<S&&>(s).first_;
|
||||
} else {
|
||||
return getv<N-1>(s.rest_);
|
||||
return getr<R,N-1>(static_cast<S&&>(s).rest_);
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
template <size_t N, bool F, bool R, class T, class... Rest>
|
||||
// Compilers disagree on decltype(auto), so these implementations avoid it at cost
|
||||
template <size_t N, bool F, bool R, class... T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
cute::enable_if_t<not cute::is_empty<cute::tuple_element_t<N, cute::type_list<T, Rest...>>>::value,
|
||||
cute::tuple_element_t<N, cute::type_list<T, Rest...>> &>
|
||||
getv(ESO<F, R, T, Rest...>& s)
|
||||
cute::conditional_t<cute::is_empty<cute::tuple_element_t<N, cute::tuple<T...>>>::value,
|
||||
cute::tuple_element_t<N, cute::tuple<T...>>,
|
||||
cute::tuple_element_t<N, cute::tuple<T...>> const&>
|
||||
getv_cr(ESO<F, R, T...> const& s) noexcept
|
||||
{
|
||||
if constexpr (N == 0) {
|
||||
return static_cast<T&>(s.first_);
|
||||
if constexpr (cute::is_empty<cute::tuple_element_t<N, cute::tuple<T...>>>::value) {
|
||||
return {};
|
||||
} else {
|
||||
return getv<N-1>(s.rest_);
|
||||
return getr<cute::tuple_element_t<N, cute::tuple<T...>> const&, N>(s);
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
template <size_t N, bool F, bool R, class T, class... Rest>
|
||||
template <size_t N, bool F, bool R, class... T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
cute::enable_if_t<not cute::is_empty<cute::tuple_element_t<N, cute::type_list<T, Rest...>>>::value,
|
||||
cute::tuple_element_t<N, cute::type_list<T, Rest...>> &&>
|
||||
getv(ESO<F, R, T, Rest...>&& s)
|
||||
cute::conditional_t<cute::is_empty<cute::tuple_element_t<N, cute::tuple<T...>>>::value,
|
||||
cute::tuple_element_t<N, cute::tuple<T...>>,
|
||||
cute::tuple_element_t<N, cute::tuple<T...>> &>
|
||||
getv_r(ESO<F, R, T...>& s) noexcept
|
||||
{
|
||||
if constexpr (N == 0) {
|
||||
return static_cast<T&&>(s.first_);
|
||||
if constexpr (cute::is_empty<cute::tuple_element_t<N, cute::tuple<T...>>>::value) {
|
||||
return {};
|
||||
} else {
|
||||
return getv<N-1>(static_cast<ESO_t<Rest...>&&>(s.rest_));
|
||||
return getr<cute::tuple_element_t<N, cute::tuple<T...>> &, N>(s);
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
template <class X, size_t N,
|
||||
bool IsFirstEmpty, bool IsRestEmpty, class First, class... Rest>
|
||||
template <size_t N, bool F, bool R, class... T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
findt(ESO<IsFirstEmpty, IsRestEmpty, First, Rest...> const& t) noexcept
|
||||
cute::conditional_t<cute::is_empty<cute::tuple_element_t<N, cute::tuple<T...>>>::value,
|
||||
cute::tuple_element_t<N, cute::tuple<T...>>,
|
||||
cute::tuple_element_t<N, cute::tuple<T...>> &&>
|
||||
getv_rr(ESO<F, R, T...>&& s) noexcept
|
||||
{
|
||||
if constexpr (cute::is_same_v<X, First>) {
|
||||
return C<N>{};
|
||||
} else
|
||||
if constexpr (sizeof...(Rest) == 0) {
|
||||
return C<N+1>{};
|
||||
} else
|
||||
if constexpr (IsRestEmpty) {
|
||||
return cute::detail::findt<X, N+1>(ESO_t<Rest...>{});
|
||||
if constexpr (cute::is_empty<cute::tuple_element_t<N, cute::tuple<T...>>>::value) {
|
||||
return {};
|
||||
} else {
|
||||
return cute::detail::findt<X, N+1>(t.rest_);
|
||||
return getr<cute::tuple_element_t<N, cute::tuple<T...>> &&, N>(static_cast<ESO<F, R, T...>&&>(s));
|
||||
}
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
} // end namespace detail
|
||||
} // end namespace eso
|
||||
|
||||
template <class... T>
|
||||
struct tuple : detail::ESO_t<T...>
|
||||
struct tuple : eso::ESO_t<T...>
|
||||
{
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
tuple() {}
|
||||
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
tuple(T const&... t) : detail::ESO_t<T...>(t...) {}
|
||||
tuple(T const&... t) : eso::ESO_t<T...>(t...) {}
|
||||
};
|
||||
|
||||
template <>
|
||||
struct tuple<> {};
|
||||
|
||||
//
|
||||
// make_tuple (value-based implementation)
|
||||
//
|
||||
|
||||
template <class... T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
tuple<T...>
|
||||
make_tuple(T const&... t)
|
||||
{
|
||||
return {t...};
|
||||
}
|
||||
|
||||
// Returns the element in the ith position of the tuple
|
||||
template <size_t I, class... T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -226,7 +227,7 @@ decltype(auto)
|
||||
get(tuple<T...> const& t) noexcept
|
||||
{
|
||||
static_assert(I < sizeof...(T), "Index out of range");
|
||||
return detail::getv<I>(t);
|
||||
return eso::getv_cr<I>(t);
|
||||
}
|
||||
|
||||
template <size_t I, class... T>
|
||||
@@ -235,7 +236,7 @@ decltype(auto)
|
||||
get(tuple<T...>& t) noexcept
|
||||
{
|
||||
static_assert(I < sizeof...(T), "Index out of range");
|
||||
return detail::getv<I>(t);
|
||||
return eso::getv_r<I>(t);
|
||||
}
|
||||
|
||||
template <size_t I, class... T>
|
||||
@@ -244,22 +245,22 @@ decltype(auto)
|
||||
get(tuple<T...>&& t) noexcept
|
||||
{
|
||||
static_assert(I < sizeof...(T), "Index out of range");
|
||||
return detail::getv<I>(static_cast<detail::ESO_t<T...>&&>(t));
|
||||
return eso::getv_rr<I>(static_cast<eso::ESO_t<T...>&&>(t));
|
||||
}
|
||||
|
||||
// Returns the position of type X (as a static integer) in the tuple
|
||||
// type's argument list. X must be unique in the argument list.
|
||||
// Returns the first position of type X (as a static integer) in the tuple
|
||||
// type's argument list.
|
||||
template <class X, class... T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
find(tuple<T...> const& t) noexcept
|
||||
find(tuple<T...> const&) noexcept
|
||||
{
|
||||
return detail::findt<X, 0>(t);
|
||||
return cute::C<find_true_v<cute::is_same_v<X,T>...>>{};
|
||||
}
|
||||
|
||||
//
|
||||
// Custom is_tuple trait simply checks the existence of tuple_size
|
||||
// and assumes std::get<I>(.), std::tuple_element<I,.>
|
||||
// and assumes get<I>(.), tuple_element<I,.>
|
||||
//
|
||||
namespace detail {
|
||||
|
||||
@@ -273,19 +274,7 @@ template <class T>
|
||||
struct is_tuple : decltype(detail::has_tuple_size((T*)0)) {};
|
||||
|
||||
template <class T>
|
||||
constexpr bool is_tuple_v = cute::is_tuple<T>::value;
|
||||
|
||||
//
|
||||
// make_tuple (value-based implementation)
|
||||
//
|
||||
|
||||
template <class... T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
tuple<T...>
|
||||
make_tuple(T const&... t)
|
||||
{
|
||||
return {t...};
|
||||
}
|
||||
static constexpr bool is_tuple_v = cute::is_tuple<T>::value;
|
||||
|
||||
//
|
||||
// tuple_cat concatenates multiple cute::tuple into a single cute::tuple,
|
||||
|
||||
@@ -31,6 +31,7 @@
|
||||
#pragma once
|
||||
|
||||
#include <cute/config.hpp> // CUTE_HOST_DEVICE, CUTE_STL_NAMESPACE
|
||||
#include <cute/util/type_traits.hpp>
|
||||
|
||||
namespace cute
|
||||
{
|
||||
@@ -39,11 +40,35 @@ template <class... T>
|
||||
struct type_list {};
|
||||
|
||||
// get<I> for type_list<T...>
|
||||
// requires tuple_element_t<I,type_list<T...>> to have std::is_default_constructible
|
||||
// Get an instance of the Ith type in the pack T...
|
||||
// Requires tuple_element_t<I,type_list<T...>> to have std::is_default_constructible
|
||||
template <size_t I, class... T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
CUTE_STL_NAMESPACE::tuple_element_t<I, type_list<T...>>
|
||||
get(type_list<T...> const& t) noexcept {
|
||||
get(type_list<T...> const&) noexcept {
|
||||
return {};
|
||||
}
|
||||
|
||||
// Find the index of the first true in the pack B...
|
||||
template <bool... B>
|
||||
struct find_true {
|
||||
CUTE_HOST_DEVICE static constexpr size_t find() {
|
||||
size_t i = 0;
|
||||
(void) ((B ? true : (++i, false)) || ...);
|
||||
return i;
|
||||
}
|
||||
static constexpr size_t value = find();
|
||||
};
|
||||
|
||||
template <bool... B>
|
||||
static constexpr size_t find_true_v = find_true<B...>::value;
|
||||
|
||||
// find<X> for type_list<T...>
|
||||
// Finds the first position of type X (as a static integer) in the T... pack
|
||||
template <class X, class... T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
CUTE_STL_NAMESPACE::integral_constant<size_t, find_true_v<cute::is_same_v<X,T>...>>
|
||||
find(type_list<T...> const&) noexcept {
|
||||
return {};
|
||||
}
|
||||
|
||||
@@ -69,9 +94,8 @@ struct tuple_size<cute::type_list<T...>>
|
||||
|
||||
template <size_t I, class... T>
|
||||
struct tuple_element<I, cute::type_list<T...>>
|
||||
{
|
||||
using type = typename CUTE_STL_NAMESPACE::tuple_element<I, CUTE_STL_NAMESPACE::tuple<T...>>::type;
|
||||
};
|
||||
: CUTE_STL_NAMESPACE::tuple_element<I, CUTE_STL_NAMESPACE::tuple<T...>>
|
||||
{};
|
||||
|
||||
} // end namespace std
|
||||
|
||||
@@ -94,9 +118,8 @@ struct tuple_size<cute::type_list<T...>>
|
||||
|
||||
template <size_t I, class... T>
|
||||
struct tuple_element<I, cute::type_list<T...>>
|
||||
{
|
||||
using type = typename CUTE_STL_NAMESPACE::tuple_element<I, CUTE_STL_NAMESPACE::tuple<T...>>::type;
|
||||
};
|
||||
: CUTE_STL_NAMESPACE::tuple_element<I, CUTE_STL_NAMESPACE::tuple<T...>>
|
||||
{};
|
||||
|
||||
} // end namespace std
|
||||
#endif // CUTE_STL_NAMESPACE_IS_CUDA_STD
|
||||
|
||||
Reference in New Issue
Block a user