v4.1 release
This commit is contained in:
@@ -42,22 +42,15 @@ namespace cute
|
||||
{
|
||||
|
||||
template <class... T>
|
||||
struct ArithmeticTuple : tuple<T...>
|
||||
{
|
||||
template <class... U>
|
||||
struct ArithmeticTuple : public tuple<T...> {
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
ArithmeticTuple(ArithmeticTuple<U...> const& u)
|
||||
: tuple<T...>(static_cast<tuple<U...> const&>(u)) {}
|
||||
ArithmeticTuple() : tuple<T...>() {}
|
||||
|
||||
template <class... U>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
ArithmeticTuple(tuple<U...> const& u)
|
||||
: tuple<T...>(u) {}
|
||||
ArithmeticTuple(tuple<T...> const& t) : tuple<T...>(t) {}
|
||||
|
||||
template <class... U>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
ArithmeticTuple(U const&... u)
|
||||
: tuple<T...>(u...) {}
|
||||
ArithmeticTuple(T const&... t) : tuple<T...>(t...) {}
|
||||
};
|
||||
|
||||
template <class... T>
|
||||
@@ -147,12 +140,12 @@ operator-(ArithmeticTuple<T...> const& t) {
|
||||
}
|
||||
|
||||
//
|
||||
// Special cases
|
||||
// Special cases for C<0>
|
||||
//
|
||||
|
||||
template <auto t, class... U>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
ArithmeticTuple<U...> const&
|
||||
ArithmeticTuple<U...>
|
||||
operator+(C<t>, ArithmeticTuple<U...> const& u) {
|
||||
static_assert(t == 0, "Arithmetic tuple op+ error!");
|
||||
return u;
|
||||
@@ -160,7 +153,7 @@ operator+(C<t>, ArithmeticTuple<U...> const& u) {
|
||||
|
||||
template <class... T, auto u>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
ArithmeticTuple<T...> const&
|
||||
ArithmeticTuple<T...>
|
||||
operator+(ArithmeticTuple<T...> const& t, C<u>) {
|
||||
static_assert(u == 0, "Arithmetic tuple op+ error!");
|
||||
return t;
|
||||
@@ -168,7 +161,7 @@ operator+(ArithmeticTuple<T...> const& t, C<u>) {
|
||||
|
||||
template <auto t, class... U>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
ArithmeticTuple<U...> const&
|
||||
ArithmeticTuple<U...>
|
||||
operator-(C<t>, ArithmeticTuple<U...> const& u) {
|
||||
static_assert(t == 0, "Arithmetic tuple op- error!");
|
||||
return -u;
|
||||
@@ -176,7 +169,7 @@ operator-(C<t>, ArithmeticTuple<U...> const& u) {
|
||||
|
||||
template <class... T, auto u>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
ArithmeticTuple<T...> const&
|
||||
ArithmeticTuple<T...>
|
||||
operator-(ArithmeticTuple<T...> const& t, C<u>) {
|
||||
static_assert(u == 0, "Arithmetic tuple op- error!");
|
||||
return t;
|
||||
@@ -212,27 +205,20 @@ struct ArithmeticTupleIterator
|
||||
}
|
||||
};
|
||||
|
||||
template <class Tuple>
|
||||
template <class... Ts>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
make_inttuple_iter(Tuple const& t) {
|
||||
return ArithmeticTupleIterator(as_arithmetic_tuple(t));
|
||||
}
|
||||
|
||||
template <class T0, class T1, class... Ts>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
make_inttuple_iter(T0 const& t0, T1 const& t1, Ts const&... ts) {
|
||||
return make_inttuple_iter(cute::make_tuple(t0, t1, ts...));
|
||||
make_inttuple_iter(Ts const&... ts) {
|
||||
return ArithmeticTupleIterator(as_arithmetic_tuple(ts...));
|
||||
}
|
||||
|
||||
//
|
||||
// ArithmeticTuple "basis" elements
|
||||
// A ScaledBasis<T,N> is a (at least) rank-N+1 ArithmeticTuple:
|
||||
// A ScaledBasis<T,Ns...> is a (at least) rank-N+1 ArithmeticTuple:
|
||||
// (_0,_0,...,T,_0,...)
|
||||
// with value T in the Nth mode
|
||||
|
||||
template <class T, int N>
|
||||
template <class T, int... Ns>
|
||||
struct ScaledBasis : private tuple<T>
|
||||
{
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
@@ -243,40 +229,61 @@ struct ScaledBasis : private tuple<T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
decltype(auto) value() const { return get<0>(static_cast<tuple<T> const&>(*this)); }
|
||||
|
||||
// Deprecated: Get the first hierarchical mode in this basis.
|
||||
CUTE_HOST_DEVICE static constexpr
|
||||
auto mode() { return Int<N>{}; }
|
||||
auto mode() { return get<0>(int_sequence<Ns...>{}); }
|
||||
};
|
||||
|
||||
// Ensure flat representation
|
||||
template <class T, int... Ms, int... Ns>
|
||||
struct ScaledBasis<ScaledBasis<T, Ms...>, Ns...> : ScaledBasis<T, Ns..., Ms...> {};
|
||||
|
||||
template <class T>
|
||||
struct is_scaled_basis : false_type {};
|
||||
template <class T, int N>
|
||||
struct is_scaled_basis<ScaledBasis<T,N>> : true_type {};
|
||||
template <class T, int... Ns>
|
||||
struct is_scaled_basis<ScaledBasis<T,Ns...>> : true_type {};
|
||||
|
||||
template <class T, int N>
|
||||
struct is_integral<ScaledBasis<T,N>> : true_type {};
|
||||
template <class T, int... Ns>
|
||||
struct is_integral<ScaledBasis<T,Ns...>> : true_type {};
|
||||
|
||||
// Get the scalar T out of a ScaledBasis
|
||||
template <class SB>
|
||||
CUTE_HOST_DEVICE constexpr auto
|
||||
basis_value(SB const& e)
|
||||
// Shortcuts
|
||||
// E<> := _1
|
||||
// E<0> := (_1,_0,_0,...)
|
||||
// E<1> := (_0,_1,_0,...)
|
||||
// E<0,0> := ((_1,_0,_0,...),_0,_0,...)
|
||||
// E<0,1> := ((_0,_1,_0,...),_0,_0,...)
|
||||
// E<1,0> := (_0,(_1,_0,_0,...),_0,...)
|
||||
// E<1,1> := (_0,(_0,_1,_0,...),_0,...)
|
||||
template <int... Ns>
|
||||
using E = ScaledBasis<Int<1>,Ns...>;
|
||||
|
||||
// Apply the Ns... pack to another Tuple
|
||||
template <class T, class Tuple>
|
||||
CUTE_HOST_DEVICE decltype(auto)
|
||||
basis_get(T const&, Tuple&& t)
|
||||
{
|
||||
if constexpr (is_scaled_basis<SB>::value) {
|
||||
return basis_value(e.value());
|
||||
return static_cast<Tuple&&>(t);
|
||||
}
|
||||
|
||||
template <class T, int... Ns, class Tuple>
|
||||
CUTE_HOST_DEVICE decltype(auto)
|
||||
basis_get(ScaledBasis<T,Ns...> const&, Tuple&& t)
|
||||
{
|
||||
if constexpr (sizeof...(Ns) == 0) {
|
||||
return static_cast<Tuple&&>(t);
|
||||
} else {
|
||||
return e;
|
||||
return get<Ns...>(static_cast<Tuple&&>(t));
|
||||
}
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
// Apply the N... pack to another Tuple
|
||||
template <class SB, class Tuple>
|
||||
template <class T>
|
||||
CUTE_HOST_DEVICE decltype(auto)
|
||||
basis_get(SB const& e, Tuple&& t)
|
||||
{
|
||||
if constexpr (is_scaled_basis<SB>::value) {
|
||||
return basis_get(e.value(), get<SB::mode()>(static_cast<Tuple&&>(t)));
|
||||
basis_value(T const& e) {
|
||||
if constexpr (is_scaled_basis<T>::value) {
|
||||
return e.value();
|
||||
} else {
|
||||
return static_cast<Tuple&&>(t);
|
||||
return e;
|
||||
}
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
@@ -294,65 +301,34 @@ to_atuple_i(T const& t, seq<I...>) {
|
||||
|
||||
// Turn a ScaledBases<T,N> into a rank-N+1 ArithmeticTuple
|
||||
// with N prefix 0s: (_0,_0,...N...,_0,T)
|
||||
template <class T, int N>
|
||||
template <class T>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
as_arithmetic_tuple(ScaledBasis<T,N> const& t) {
|
||||
return detail::to_atuple_i(as_arithmetic_tuple(t.value()), make_seq<N>{});
|
||||
as_arithmetic_tuple(ScaledBasis<T> const& t) {
|
||||
return t.value();
|
||||
}
|
||||
|
||||
namespace detail {
|
||||
template <class T, int N, int... Ns>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
as_arithmetic_tuple(ScaledBasis<T,N,Ns...> const& t) {
|
||||
return detail::to_atuple_i(as_arithmetic_tuple(ScaledBasis<T,Ns...>{t.value()}), make_seq<N>{});
|
||||
}
|
||||
|
||||
template <int... Ns>
|
||||
struct Basis;
|
||||
|
||||
template <>
|
||||
struct Basis<> {
|
||||
using type = Int<1>;
|
||||
};
|
||||
|
||||
template <int N, int... Ns>
|
||||
struct Basis<N,Ns...> {
|
||||
using type = ScaledBasis<typename Basis<Ns...>::type, N>;
|
||||
};
|
||||
|
||||
} // end namespace detail
|
||||
|
||||
// Shortcut for writing ScaledBasis<ScaledBasis<ScaledBasis<Int<1>, N0>, N1>, ...>
|
||||
// E<> := _1
|
||||
// E<0> := (_1,_0,_0,...)
|
||||
// E<1> := (_0,_1,_0,...)
|
||||
// E<0,0> := ((_1,_0,_0,...),_0,_0,...)
|
||||
// E<0,1> := ((_0,_1,_0,...),_0,_0,...)
|
||||
// E<1,0> := (_0,(_1,_0,_0,...),_0,...)
|
||||
// E<1,1> := (_0,(_0,_1,_0,...),_0,...)
|
||||
template <int... N>
|
||||
using E = typename detail::Basis<N...>::type;
|
||||
|
||||
template <class Shape>
|
||||
template <int... Ns, class Shape>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
make_basis_like(Shape const& shape)
|
||||
{
|
||||
if constexpr (is_integral<Shape>::value) {
|
||||
return Int<1>{};
|
||||
} else {
|
||||
// Generate bases for each rank of shape
|
||||
if constexpr (is_tuple<Shape>::value) {
|
||||
// Generate bases for each mode of shape
|
||||
return transform(tuple_seq<Shape>{}, shape, [](auto I, auto si) {
|
||||
// Generate bases for each rank of si and add an i on front
|
||||
using I_type = decltype(I);
|
||||
return transform_leaf(make_basis_like(si), [](auto e) {
|
||||
// MSVC has trouble capturing variables as constexpr,
|
||||
// so that they can be used as template arguments.
|
||||
// This is exactly what the code needs to do with i, unfortunately.
|
||||
// The work-around is to define i inside the inner lambda,
|
||||
// by using just the type from the enclosing scope.
|
||||
constexpr int i = I_type::value;
|
||||
return ScaledBasis<decltype(e), i>{};
|
||||
});
|
||||
// Generate bases for each si and add an i on end
|
||||
return make_basis_like<Ns...,decltype(I)::value>(si);
|
||||
});
|
||||
} else {
|
||||
return E<Ns...>{};
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
@@ -360,109 +336,124 @@ make_basis_like(Shape const& shape)
|
||||
// Arithmetic
|
||||
//
|
||||
|
||||
template <class T, int M, class U>
|
||||
template <class T, int... Ns, class U>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
safe_div(ScaledBasis<T,M> const& b, U const& u)
|
||||
safe_div(ScaledBasis<T,Ns...> const& b, U const& u)
|
||||
{
|
||||
auto t = safe_div(b.value(), u);
|
||||
return ScaledBasis<decltype(t),M>{t};
|
||||
return ScaledBasis<decltype(t),Ns...>{t};
|
||||
}
|
||||
|
||||
template <class T, int M, class U>
|
||||
template <class T, int... Ns, class U>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
ceil_div(ScaledBasis<T,M> const& b, U const& u)
|
||||
ceil_div(ScaledBasis<T,Ns...> const& b, U const& u)
|
||||
{
|
||||
auto t = ceil_div(b.value(), u);
|
||||
return ScaledBasis<decltype(t),M>{t};
|
||||
return ScaledBasis<decltype(t),Ns...>{t};
|
||||
}
|
||||
|
||||
template <class T, int N>
|
||||
template <class T, int... Ns>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
abs(ScaledBasis<T,N> const& e)
|
||||
abs(ScaledBasis<T,Ns...> const& e)
|
||||
{
|
||||
auto t = abs(e.value());
|
||||
return ScaledBasis<decltype(t),N>{t};
|
||||
return ScaledBasis<decltype(t),Ns...>{t};
|
||||
}
|
||||
|
||||
// Equality
|
||||
template <class T, int N, class U, int M>
|
||||
template <class T, int... Ns, class U, int... Ms>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator==(ScaledBasis<T,N> const& t, ScaledBasis<U,M> const& u) {
|
||||
return bool_constant<M == N>{} && t.value() == u.value();
|
||||
operator==(ScaledBasis<T,Ns...> const& t, ScaledBasis<U,Ms...> const& u) {
|
||||
if constexpr (sizeof...(Ns) == sizeof...(Ms)) {
|
||||
return bool_constant<((Ns == Ms) && ...)>{} && t.value() == u.value();
|
||||
} else {
|
||||
return false_type{};
|
||||
}
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
// Not equal to anything else
|
||||
template <class T, int N, class U>
|
||||
template <class T, int... Ns, class U>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
false_type
|
||||
operator==(ScaledBasis<T,N> const&, U const&) {
|
||||
operator==(ScaledBasis<T,Ns...> const&, U const&) {
|
||||
return {};
|
||||
}
|
||||
|
||||
template <class T, class U, int M>
|
||||
template <class T, class U, int... Ms>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
false_type
|
||||
operator==(T const&, ScaledBasis<U,M> const&) {
|
||||
operator==(T const&, ScaledBasis<U,Ms...> const&) {
|
||||
return {};
|
||||
}
|
||||
|
||||
// Multiplication
|
||||
template <class A, class T, int N>
|
||||
template <class A, class T, int... Ns>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator*(A const& a, ScaledBasis<T,N> const& e) {
|
||||
operator*(A const& a, ScaledBasis<T,Ns...> const& e) {
|
||||
auto r = a * e.value();
|
||||
return ScaledBasis<decltype(r),N>{r};
|
||||
return ScaledBasis<decltype(r),Ns...>{r};
|
||||
}
|
||||
|
||||
template <class T, int N, class B>
|
||||
template <class T, int... Ns, class B>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator*(ScaledBasis<T,N> const& e, B const& b) {
|
||||
operator*(ScaledBasis<T,Ns...> const& e, B const& b) {
|
||||
auto r = e.value() * b;
|
||||
return ScaledBasis<decltype(r),N>{r};
|
||||
return ScaledBasis<decltype(r),Ns...>{r};
|
||||
}
|
||||
|
||||
// Addition
|
||||
template <class T, int N, class U, int M>
|
||||
template <class T, int... Ns, class U, int... Ms>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator+(ScaledBasis<T,N> const& t, ScaledBasis<U,M> const& u) {
|
||||
operator+(ScaledBasis<T,Ns...> const& t, ScaledBasis<U,Ms...> const& u) {
|
||||
return as_arithmetic_tuple(t) + as_arithmetic_tuple(u);
|
||||
}
|
||||
|
||||
template <class T, int N, class... U>
|
||||
template <class T, int... Ns, class... U>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator+(ScaledBasis<T,N> const& t, ArithmeticTuple<U...> const& u) {
|
||||
operator+(ScaledBasis<T,Ns...> const& t, ArithmeticTuple<U...> const& u) {
|
||||
return as_arithmetic_tuple(t) + u;
|
||||
}
|
||||
|
||||
template <class... T, class U, int M>
|
||||
template <class... T, class U, int... Ms>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator+(ArithmeticTuple<T...> const& t, ScaledBasis<U,M> const& u) {
|
||||
operator+(ArithmeticTuple<T...> const& t, ScaledBasis<U,Ms...> const& u) {
|
||||
return t + as_arithmetic_tuple(u);
|
||||
}
|
||||
|
||||
template <auto t, class U, int M>
|
||||
template <auto t, class U, int... Ms>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator+(C<t>, ScaledBasis<U,M> const& u) {
|
||||
static_assert(t == 0, "ScaledBasis op+ error!");
|
||||
return u;
|
||||
operator+(C<t>, ScaledBasis<U,Ms...> const& u) {
|
||||
if constexpr (sizeof...(Ms) == 0) {
|
||||
return C<t>{} + u.value();
|
||||
} else {
|
||||
static_assert(t == 0, "ScaledBasis op+ error!");
|
||||
return u;
|
||||
}
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
template <class T, int N, auto u>
|
||||
template <class T, int... Ns, auto u>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator+(ScaledBasis<T,N> const& t, C<u>) {
|
||||
static_assert(u == 0, "ScaledBasis op+ error!");
|
||||
return t;
|
||||
operator+(ScaledBasis<T,Ns...> const& t, C<u>) {
|
||||
if constexpr (sizeof...(Ns) == 0) {
|
||||
return t.value() + C<u>{};
|
||||
} else {
|
||||
static_assert(u == 0, "ScaledBasis op+ error!");
|
||||
return t;
|
||||
}
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
}
|
||||
|
||||
//
|
||||
@@ -475,10 +466,10 @@ CUTE_HOST_DEVICE void print(ArithmeticTupleIterator<ArithTuple> const& iter)
|
||||
printf("ArithTuple"); print(iter.coord_);
|
||||
}
|
||||
|
||||
template <class T, int N>
|
||||
CUTE_HOST_DEVICE void print(ScaledBasis<T,N> const& e)
|
||||
template <class T, int... Ns>
|
||||
CUTE_HOST_DEVICE void print(ScaledBasis<T,Ns...> const& e)
|
||||
{
|
||||
print(e.value()); printf("@%d", N);
|
||||
print(e.value()); (void(printf("@%d", Ns)), ...);
|
||||
}
|
||||
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
@@ -488,10 +479,11 @@ CUTE_HOST std::ostream& operator<<(std::ostream& os, ArithmeticTupleIterator<Ari
|
||||
return os << "ArithTuple" << iter.coord_;
|
||||
}
|
||||
|
||||
template <class T, int N>
|
||||
CUTE_HOST std::ostream& operator<<(std::ostream& os, ScaledBasis<T,N> const& e)
|
||||
template <class T, int... Ns>
|
||||
CUTE_HOST std::ostream& operator<<(std::ostream& os, ScaledBasis<T,Ns...> const& e)
|
||||
{
|
||||
return os << e.value() << "@" << N;
|
||||
os << e.value(); (void(os << "@" << Ns), ...);
|
||||
return os;
|
||||
}
|
||||
#endif
|
||||
|
||||
|
||||
@@ -47,8 +47,9 @@ namespace cute
|
||||
// Signed integers
|
||||
//
|
||||
|
||||
using int2_t = cutlass::int2b_t;
|
||||
using int4_t = cutlass::int4b_t;
|
||||
using int2_t = cutlass::int2b_t;
|
||||
using int4_t = cutlass::int4b_t;
|
||||
using int6_t = cutlass::int6b_t;
|
||||
using CUTE_STL_NAMESPACE::int8_t;
|
||||
using CUTE_STL_NAMESPACE::int16_t;
|
||||
using CUTE_STL_NAMESPACE::int32_t;
|
||||
@@ -75,10 +76,10 @@ using int_byte_t = typename int_byte<N>::type;
|
||||
// Unsigned integers
|
||||
//
|
||||
|
||||
using uint1_t = cutlass::uint1b_t;
|
||||
using uint2_t = cutlass::uint2b_t;
|
||||
using uint4_t = cutlass::uint4b_t;
|
||||
using uint6_t = cutlass::uint6b_t;
|
||||
using uint1_t = cutlass::uint1b_t;
|
||||
using uint2_t = cutlass::uint2b_t;
|
||||
using uint4_t = cutlass::uint4b_t;
|
||||
using uint6_t = cutlass::uint6b_t;
|
||||
using CUTE_STL_NAMESPACE::uint8_t;
|
||||
using CUTE_STL_NAMESPACE::uint16_t;
|
||||
using CUTE_STL_NAMESPACE::uint32_t;
|
||||
@@ -88,7 +89,7 @@ template <int N> struct uint_bit;
|
||||
template <> struct uint_bit< 1> { using type = uint1_t; };
|
||||
template <> struct uint_bit< 2> { using type = uint2_t; };
|
||||
template <> struct uint_bit< 4> { using type = uint4_t; };
|
||||
template <> struct uint_bit< 6> { using type = uint6_t; };
|
||||
template <> struct uint_bit< 6> { using type = uint6_t; };
|
||||
template <> struct uint_bit< 8> { using type = uint8_t; };
|
||||
template <> struct uint_bit< 16> { using type = uint16_t; };
|
||||
template <> struct uint_bit< 32> { using type = uint32_t; };
|
||||
|
||||
@@ -38,10 +38,19 @@
|
||||
|
||||
namespace cute {
|
||||
|
||||
template <typename T>
|
||||
struct sizeof_bits : public cutlass::sizeof_bits<T> {};
|
||||
template <class T>
|
||||
struct sizeof_bits : cutlass::sizeof_bits<T> {};
|
||||
|
||||
// DO NOT change auto to int, sizeof_bits<sparse_elem> use integral_ratio instead of int
|
||||
template <class T>
|
||||
struct sizeof_bits<T const> : sizeof_bits<T> {};
|
||||
|
||||
template <class T>
|
||||
struct sizeof_bits<T volatile> : sizeof_bits<T> {};
|
||||
|
||||
template <class T>
|
||||
struct sizeof_bits<T const volatile> : sizeof_bits<T> {};
|
||||
|
||||
// DO NOT change auto to int, sizeof_bits<sparse_elem> use integral_ratio instead of int
|
||||
template <class T>
|
||||
static constexpr auto sizeof_bits_v = sizeof_bits<T>::value;
|
||||
|
||||
@@ -53,6 +62,23 @@ using cutlass::is_subbyte;
|
||||
template <class T>
|
||||
static constexpr auto is_subbyte_v = is_subbyte<T>::value;
|
||||
|
||||
//
|
||||
// Integral
|
||||
//
|
||||
|
||||
using cutlass::bin1_t;
|
||||
using cutlass::uint1b_t;
|
||||
using cutlass::int2b_t;
|
||||
using cutlass::uint2b_t;
|
||||
using cutlass::int4b_t;
|
||||
using cutlass::uint4b_t;
|
||||
using cutlass::int6b_t;
|
||||
using cutlass::uint6b_t;
|
||||
|
||||
//
|
||||
// Floating Point
|
||||
//
|
||||
|
||||
using cutlass::half_t;
|
||||
using cutlass::bfloat16_t;
|
||||
|
||||
@@ -65,18 +91,12 @@ using cutlass::type_erased_dynamic_float8_t;
|
||||
using cutlass::float_e4m3_t;
|
||||
using cutlass::float_e5m2_t;
|
||||
|
||||
using cutlass::uint1b_t;
|
||||
using cutlass::int2b_t;
|
||||
using cutlass::uint2b_t;
|
||||
using cutlass::int4b_t;
|
||||
using cutlass::uint4b_t;
|
||||
using cutlass::bin1_t;
|
||||
|
||||
|
||||
|
||||
using cutlass::float_ue4m3_t;
|
||||
using cutlass::float_ue8m0_t;
|
||||
|
||||
using cutlass::uint6b_t;
|
||||
using cutlass::float_e2m1_t;
|
||||
using cutlass::float_e2m3_t;
|
||||
using cutlass::float_e3m2_t;
|
||||
@@ -94,8 +114,6 @@ using cutlass::detail::type_erased_dynamic_float4_unpacksmem_t;
|
||||
using cutlass::detail::type_erased_dynamic_float6_unpacksmem_t;
|
||||
};
|
||||
|
||||
|
||||
|
||||
//
|
||||
// Print utility
|
||||
//
|
||||
@@ -112,7 +130,6 @@ print(bfloat16_t a) {
|
||||
printf("%f", static_cast<float>(a));
|
||||
}
|
||||
|
||||
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print(tfloat32_t a) {
|
||||
@@ -131,6 +148,15 @@ print(float_e5m2_t a) {
|
||||
printf("%f", static_cast<float>(a));
|
||||
}
|
||||
|
||||
template <cutlass::detail::FpEncoding Encoding, class Derived>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print(cutlass::float_exmy_base<Encoding, Derived> a) {
|
||||
printf("%f", static_cast<float>(a));
|
||||
}
|
||||
|
||||
// Pretty Print utility
|
||||
|
||||
CUTE_HOST_DEVICE void
|
||||
pretty_print(bfloat16_t v) {
|
||||
printf("%*.2f", 8, float(v));
|
||||
@@ -156,26 +182,11 @@ pretty_print(float_e5m2_t t) {
|
||||
printf("%*.2f", 8, static_cast<float>(t));
|
||||
}
|
||||
|
||||
|
||||
template <
|
||||
cutlass::detail::FpEncoding Encoding,
|
||||
class Derived
|
||||
>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
print(cutlass::float_exmy_base<Encoding, Derived> a) {
|
||||
printf("%f", static_cast<float>(a));
|
||||
}
|
||||
|
||||
template <
|
||||
cutlass::detail::FpEncoding Encoding,
|
||||
class Derived
|
||||
>
|
||||
template <cutlass::detail::FpEncoding Encoding, class Derived>
|
||||
CUTE_HOST_DEVICE
|
||||
void
|
||||
pretty_print_float_exmy_base(cutlass::float_exmy_base<Encoding, Derived> t) {
|
||||
printf("%*.2f", 8, static_cast<float>(t));
|
||||
}
|
||||
|
||||
|
||||
} // namespace cute
|
||||
|
||||
Reference in New Issue
Block a user