CUTLASS 3.5.0 (#1411)
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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;
|
||||
};
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
{
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
Reference in New Issue
Block a user