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:
Yujia Zhai
2025-04-24 15:42:40 -04:00
committed by GitHub
co-authored by yuzhai Haicheng Wu
parent 8e345c5c5b
commit 331a1f5b3f
143 changed files with 18089 additions and 5935 deletions
+63 -74
View File
@@ -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 -8
View File
@@ -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