CUTLASS 3.5.0 (#1411)

This commit is contained in:
Vijay Thakkar
2024-03-19 17:51:04 -04:00
committed by GitHub
parent ffa34e7075
commit 629f4653c3
468 changed files with 48729 additions and 7252 deletions
+1 -1
View File
@@ -32,7 +32,7 @@
#include <cute/config.hpp>
#include <cute/numeric/int.hpp>
#include <cute/numeric/numeric_types.hpp>
#include <cute/numeric/math.hpp>
namespace cute
+13 -13
View File
@@ -355,7 +355,7 @@ void clear(array<T,N>& a)
a.fill(T(0));
}
template <typename T, size_t N>
template <class T, size_t N>
CUTE_HOST_DEVICE constexpr
void fill(array<T,N>& a, T const& value)
{
@@ -370,14 +370,14 @@ void swap(array<T,N>& a, array<T,N>& b)
}
/// @return A cute::array of the elements of @c t in reverse order.
template <typename T, size_t N>
CUTE_HOST_DEVICE constexpr cute::array<T, N>
reverse(cute::array<T, N> const& t) {
template <class T, size_t N>
CUTE_HOST_DEVICE constexpr
cute::array<T,N> reverse(cute::array<T,N> const& t)
{
if constexpr (N == 0u) {
return t;
}
else {
cute::array<T, N> t_r{};
} else {
cute::array<T,N> t_r{};
for (size_t k = 0; k < N; ++k) {
t_r[k] = t[N - k - 1];
}
@@ -422,7 +422,7 @@ CUTE_HOST_DEVICE constexpr
T&& get(array<T,N>&& a)
{
static_assert(I < N, "Index out of range");
return std::move(a[I]);
return cute::move(a[I]);
}
} // end namespace cute
@@ -442,12 +442,12 @@ struct tuple_element<I, cute::array<T,N>>
};
template <class T, size_t N>
struct tuple_size<const cute::array<T,N>>
struct tuple_size<cute::array<T,N> const>
: CUTE_STL_NAMESPACE::integral_constant<size_t, N>
{};
template <size_t I, class T, size_t N>
struct tuple_element<I, const cute::array<T,N>>
struct tuple_element<I, cute::array<T,N> const>
{
using type = T;
};
@@ -462,7 +462,7 @@ namespace std
template <class... _Tp>
struct tuple_size;
template<size_t _Ip, class... _Tp>
template <size_t _Ip, class... _Tp>
struct tuple_element;
#endif
@@ -478,12 +478,12 @@ struct tuple_element<I, cute::array<T,N>>
};
template <class T, size_t N>
struct tuple_size<const cute::array<T,N>>
struct tuple_size<cute::array<T,N> const>
: CUTE_STL_NAMESPACE::integral_constant<size_t, N>
{};
template <size_t I, class T, size_t N>
struct tuple_element<I, const cute::array<T,N>>
struct tuple_element<I, cute::array<T,N> const>
{
using type = T;
};
+17 -21
View File
@@ -37,29 +37,20 @@
#include <cute/config.hpp>
#include <cute/numeric/int.hpp> // sizeof_bits
#include <cute/numeric/numeric_types.hpp>
#include <cute/numeric/integral_constant.hpp>
namespace cute
{
template <class T>
struct is_subbyte {
static constexpr bool value = sizeof_bits_v<T> < 8;
};
template <class T>
constexpr bool is_subbyte_v = is_subbyte<T>::value;
//
// Underlying subbyte storage type
//
template <class T>
using subbyte_storage_type_t = conditional_t<(sizeof_bits_v<T> <= 8), uint8_t,
conditional_t<(sizeof_bits_v<T> <= 16), uint16_t,
conditional_t<(sizeof_bits_v<T> <= 32), uint32_t,
conditional_t<(sizeof_bits_v<T> <= 64), uint64_t,
conditional_t<(sizeof_bits_v<T> <= 128), uint128_t,
using subbyte_storage_type_t = conditional_t<(cute::sizeof_bits_v<T> <= 8), uint8_t,
conditional_t<(cute::sizeof_bits_v<T> <= 16), uint16_t,
conditional_t<(cute::sizeof_bits_v<T> <= 32), uint32_t,
conditional_t<(cute::sizeof_bits_v<T> <= 64), uint64_t,
conditional_t<(cute::sizeof_bits_v<T> <= 128), uint128_t,
T>>>>>;
template <class T> struct subbyte_iterator;
@@ -183,6 +174,11 @@ public:
operator element_type() const {
return get();
}
// Address
subbyte_iterator<T> operator&() const {
return {ptr_, idx_};
}
};
//
@@ -314,7 +310,7 @@ public:
CUTE_HOST_DEVICE constexpr friend
auto recast_ptr(subbyte_iterator const& x) {
using NewT = conditional_t<(is_const_v<T>), NewT_ const, NewT_>;
if constexpr (is_subbyte<NewT>::value) { // Making subbyte_iter, preserve the subbyte idx
if constexpr (cute::is_subbyte_v<NewT>) { // Making subbyte_iter, preserve the subbyte idx
return subbyte_iterator<NewT>(x.ptr_, x.idx_);
} else { // Not subbyte, assume/assert subbyte idx 0
return reinterpret_cast<NewT*>(raw_pointer_cast(x));
@@ -323,7 +319,7 @@ public:
}
CUTE_HOST_DEVICE friend void print(subbyte_iterator x) {
printf("subptr[%db](%p.%u)", int(sizeof_bits<T>::value), x.ptr_, x.idx_);
printf("subptr[%db](%p.%u)", int(sizeof_bits_v<T>), x.ptr_, x.idx_);
}
};
@@ -369,8 +365,8 @@ private:
public:
CUTE_HOST_DEVICE constexpr
array_subbyte() {}
constexpr
array_subbyte() = default;
CUTE_HOST_DEVICE constexpr
array_subbyte(array_subbyte const& x) {
@@ -562,7 +558,7 @@ CUTE_HOST_DEVICE constexpr
T&& get(array_subbyte<T,N>&& a)
{
static_assert(I < N, "Index out of range");
return std::move(a[I]);
return cute::move(a[I]);
}
} // end namespace cute
@@ -608,7 +604,7 @@ namespace std
template <class... _Tp>
struct tuple_size;
template<size_t _Ip, class... _Tp>
template <size_t _Ip, class... _Tp>
struct tuple_element;
#endif
+1 -1
View File
@@ -37,7 +37,7 @@
#include <cute/config.hpp>
#include <cute/numeric/int.hpp> // uint_bit_t
#include <cute/numeric/numeric_types.hpp> // uint_bit_t
namespace cute
{
+6 -6
View File
@@ -96,11 +96,11 @@ uint32_t&& get(dim3&& a)
{
static_assert(I < 3, "Index out of range");
if constexpr (I == 0) {
return std::move(a.x);
return cute::move(a.x);
} else if constexpr (I == 1) {
return std::move(a.y);
return cute::move(a.y);
} else if constexpr (I == 2) {
return std::move(a.z);
return cute::move(a.z);
}
CUTE_GCC_UNREACHABLE;
@@ -162,11 +162,11 @@ uint32_t&& get(uint3&& a)
{
static_assert(I < 3, "Index out of range");
if constexpr (I == 0) {
return std::move(a.x);
return cute::move(a.x);
} else if constexpr (I == 1) {
return std::move(a.y);
return cute::move(a.y);
} else if constexpr (I == 2) {
return std::move(a.z);
return cute::move(a.z);
}
CUTE_GCC_UNREACHABLE;
+10 -25
View File
@@ -126,18 +126,14 @@ CUTE_HOST_DEVICE constexpr T& getv(EBO<N, T, false>& x)
template <size_t N, class T>
CUTE_HOST_DEVICE constexpr T&& getv(EBO<N, T, false>&& x)
{ return static_cast<T&&>(x.t_); }
{ return cute::move(x.t_); }
template <class IdxSeq, class... T>
struct TupleBase;
// Base class of cute::tuple.
// It inherits from EBO<i, t> for each (i, t) in (I..., T...).
// The actual storage (for nonempty t) lives in the base classes.
// index_sequence is a way to wrap up a sequence of zero or more
// compile-time integer values in a single type.
// We only ever use index_sequence<0, 1, ..., sizeof...(T)> in practice,
// as the type alias TupleBase below indicates.
// Base class of cute::tuple binds each element to an index
// by inheriting from EBO<i, t> for each (i, t) in (I..., T...).
// The storage (for nonempty t) lives in the base classes.
template <size_t... I, class... T>
struct TupleBase<index_sequence<I...>, T...>
: EBO<I,T>...
@@ -169,11 +165,6 @@ struct TupleBase<index_sequence<I...>, T...>
//
// Inheriting from the above alias TupleBase
// causes MSVC 2022 build errors when assigning one tuple to another:
//
// illegal member initialization:
// 'TupleBase< /* template arguments */ >' is not a base or member
//
// Not using the alias or any kind of alias fixed the errors.
// In summary: this is verbose as a work-around for MSVC build errors.
template <class... T>
struct tuple : detail::TupleBase<make_index_sequence<sizeof...(T)>, T...>
@@ -365,10 +356,10 @@ tuple_cat(T0 const& t0, T1 const& t1, T2 const& t2, T3 const& t3, T4 const& t4,
return cute::make_tuple(get<I0>(t0)..., get<I1>(t1)..., get<I2>(t2)..., get<I3>(t3)..., get<I4>(t4)...);
}
template<class T0, class T1>
template <class T0, class T1>
struct tuple_cat_static;
template<class... T0s, class... T1s>
template <class... T0s, class... T1s>
struct tuple_cat_static<tuple<T0s...>, tuple<T1s...>> {
using type = tuple<T0s..., T1s...>;
};
@@ -630,11 +621,8 @@ template <class Tuple, size_t... Is>
CUTE_HOST_DEVICE void print_tuple(Tuple const& t,
index_sequence<Is...>, char s = '(', char e = ')')
{
using eat = int[];
using cute::print;
(void) eat {(print(s), 0),
(print(Is == 0 ? "" : ","), print(get<Is>(t)), 0)...,
(print(e), 0)};
((void(print(Is == 0 ? s : ',')), void(print(get<Is>(t)))), ...); print(e);
}
#if !defined(__CUDACC_RTC__)
@@ -642,11 +630,8 @@ template <class Tuple, std::size_t... Is>
CUTE_HOST std::ostream& print_tuple_os(std::ostream& os, Tuple const& t,
index_sequence<Is...>, char s = '(', char e = ')')
{
using eat = int[];
(void) eat {(void(os << s), 0),
(void(os << (Is == 0 ? "" : ",") << get<Is>(t)), 0)...,
(void(os << e), 0)};
return os;
(void(os << (Is == 0 ? s : ',') << get<Is>(t)), ...);
return os << e;
}
#endif // !defined(__CUDACC_RTC__)
@@ -707,7 +692,7 @@ namespace std
template <class... _Tp>
struct tuple_size;
template<size_t _Ip, class... _Tp>
template <size_t _Ip, class... _Tp>
struct tuple_element;
#endif
+1 -1
View File
@@ -108,7 +108,7 @@ namespace std
template <class... _Tp>
struct tuple_size;
template<size_t _Ip, class... _Tp>
template <size_t _Ip, class... _Tp>
struct tuple_element;
#endif