CUTLASS 3.2.1 (#1113)
* Updates for 3.2.1 release. * Minor fix in gemm op profiler for raster order. * Add scheduler mapping for raster order in the kernels.
This commit is contained in:
+85
-86
@@ -218,41 +218,40 @@ recast(Swizzle<B,M,S> const& swizzle)
|
||||
// consumed and which bits are free. Furthermore, it is useful to know whether
|
||||
// each of these bits is known statically or dynamically.
|
||||
|
||||
// MixedBits is an integer class where some bits are known statically and some
|
||||
// bits are known dynamically. These sets of bits are disjoint and it is known
|
||||
// statically which bits are known dynamically.
|
||||
// MixedBits is an 32-bit unsigned integer class where some bits are known statically
|
||||
// and some bits are known dynamically. These sets of bits are disjoint and it is
|
||||
// known statically which bits are known dynamically.
|
||||
|
||||
// MixedBits can only be manipulated through bitwise operations
|
||||
|
||||
// Abstract value: StaticInt | (dynamic_int_ & StaticFlags)
|
||||
template <uint32_t StaticInt = 0,
|
||||
class DynamicType = uint32_t,
|
||||
uint32_t StaticFlags = 0> // 0: static, 1: dynamic
|
||||
template <uint32_t StaticInt,
|
||||
uint32_t StaticFlags> // 0: static, 1: dynamic
|
||||
struct MixedBits
|
||||
{
|
||||
// Representation invariants
|
||||
static_assert(StaticFlags != 0, "Should be at least one dynamic bit in MixedBits.");
|
||||
static_assert((StaticInt & StaticFlags) == 0, "No static/dynamic overlap allowed in MixedBits.");
|
||||
// assert((dynamic_int_ & ~F) == 0);
|
||||
|
||||
DynamicType dynamic_int_;
|
||||
uint32_t dynamic_int_;
|
||||
// assert((dynamic_int_ & ~StaticFlags) == 0);
|
||||
|
||||
CUTE_HOST_DEVICE constexpr operator uint32_t() const noexcept { return StaticInt | dynamic_int_; }
|
||||
};
|
||||
|
||||
template <class S, S s, class DynamicType, class F, F f>
|
||||
// Return a value representing (C<s>{} | (d & C<f>)) potentially using MixedBits to track s and f.
|
||||
// This maker does allow ((s & f) != 0) and enforces the MixedBits invariant before creation.
|
||||
template <auto s, class DynamicType, auto f>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
make_mixed_bits(constant<S,s> const&, DynamicType const& d, constant<F,f> const&)
|
||||
make_mixed_bits(C<s>, DynamicType const& d, C<f>)
|
||||
{
|
||||
static_assert(is_integral<DynamicType>::value);
|
||||
if constexpr (is_static<DynamicType>::value) {
|
||||
static_assert((s & DynamicType::value & f) == 0, "No static/dynamic overlap allowed.");
|
||||
return constant<S,s>{} | (d & constant<F,f>{}); // Just return a static int
|
||||
} else if constexpr (f == 0) {
|
||||
return constant<S,s>{}; // Just return a static int
|
||||
constexpr uint32_t new_f = uint32_t(f) & ~uint32_t(s); // StaticBits take precedence, M<0,f>{d} | C<s>{}
|
||||
if constexpr (new_f == 0 || is_static<DynamicType>::value) {
|
||||
return C<s>{} | (d & C<new_f>{}); // Just return a static int
|
||||
} else {
|
||||
return MixedBits<s, DynamicType, f>{d & f}; // MixedBits
|
||||
return MixedBits<s, new_f>{uint32_t(d) & new_f}; // MixedBits
|
||||
}
|
||||
|
||||
CUTE_GCC_UNREACHABLE;
|
||||
@@ -263,28 +262,28 @@ make_mixed_bits(constant<S,s> const&, DynamicType const& d, constant<F,f> const&
|
||||
//
|
||||
|
||||
// Equality
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator==(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const&)
|
||||
operator==(MixedBits<S0,F0> const& m, C<S1>)
|
||||
{
|
||||
return (S0 == (S1 & ~F0)) && (m.dynamic_int_ == (S1 & F0));
|
||||
return (S0 == (uint32_t(S1) & ~F0)) && (m.dynamic_int_ == (uint32_t(S1) & F0));
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator==(constant<TS1,S1> const& s, MixedBits<S0,D0,F0> const& m)
|
||||
operator==(C<S1> s, MixedBits<S0,F0> const& m)
|
||||
{
|
||||
return m == s;
|
||||
}
|
||||
|
||||
// Bitwise AND
|
||||
template <uint32_t S0, class D0, uint32_t F0,
|
||||
uint32_t S1, class D1, uint32_t F1>
|
||||
template <uint32_t S0, uint32_t F0,
|
||||
uint32_t S1, uint32_t F1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator&(MixedBits<S0,D0,F0> const& m0, MixedBits<S1,D1,F1> const& m1)
|
||||
operator&(MixedBits<S0,F0> const& m0, MixedBits<S1,F1> const& m1)
|
||||
{
|
||||
// Truth table for (S0,D0,F0) & (S1,D1,F1) -> (S,D,F)
|
||||
// S0D0F0 | 0X0 | 001 | 011 | 1X0 |
|
||||
@@ -294,36 +293,36 @@ operator&(MixedBits<S0,D0,F0> const& m0, MixedBits<S1,D1,F1> const& m1)
|
||||
// 011 | 0X0 | 001 | 011 | 011 |
|
||||
// 1X0 | 0X0 | 001 | 011 | 1X0 |
|
||||
|
||||
return make_mixed_bits(constant<uint32_t,S0 & S1>{},
|
||||
return make_mixed_bits(C<S0 & S1>{},
|
||||
//(S0 | m0.dynamic_int_) & (S1 | m1.dynamic_int_),
|
||||
((S1 & F0) & m0.dynamic_int_) | ((S0 & F1) & m1.dynamic_int_) | (m0.dynamic_int_ & m1.dynamic_int_),
|
||||
constant<uint32_t,(S1 & F0) | (S0 & F1) | (F0 & F1)>{});
|
||||
C<(S1 & F0) | (S0 & F1) | (F0 & F1)>{});
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator&(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const&)
|
||||
operator&(MixedBits<S0,F0> const& m, C<S1>)
|
||||
{
|
||||
return make_mixed_bits(constant<uint32_t,S0 & S1>{},
|
||||
return make_mixed_bits(C<S0 & uint32_t(S1)>{},
|
||||
m.dynamic_int_,
|
||||
constant<uint32_t,S1 & F0>{});
|
||||
C<F0 & uint32_t(S1)>{});
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator&(constant<TS1,S1> const& s, MixedBits<S0,D0,F0> const& m)
|
||||
operator&(C<S1> s, MixedBits<S0,F0> const& m)
|
||||
{
|
||||
return m & s;
|
||||
}
|
||||
|
||||
// Bitwise OR
|
||||
template <uint32_t S0, class D0, uint32_t F0,
|
||||
uint32_t S1, class D1, uint32_t F1>
|
||||
template <uint32_t S0, uint32_t F0,
|
||||
uint32_t S1, uint32_t F1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator|(MixedBits<S0,D0,F0> const& m0, MixedBits<S1,D1,F1> const& m1)
|
||||
operator|(MixedBits<S0,F0> const& m0, MixedBits<S1,F1> const& m1)
|
||||
{
|
||||
// Truth table for (S0,D0,F0) | (S1,D1,F1) -> (S,D,F)
|
||||
// S0D0F0 | 0X0 | 001 | 011 | 1X0 |
|
||||
@@ -333,35 +332,35 @@ operator|(MixedBits<S0,D0,F0> const& m0, MixedBits<S1,D1,F1> const& m1)
|
||||
// 011 | 011 | 011 | 011 | 1X0 |
|
||||
// 1X0 | 1X0 | 1X0 | 1X0 | 1X0 |
|
||||
|
||||
return make_mixed_bits(constant<uint32_t,S0 | S1>{},
|
||||
return make_mixed_bits(C<S0 | S1>{},
|
||||
((~S1 & F0) & m0.dynamic_int_) | ((~S0 & F1) & m1.dynamic_int_),
|
||||
constant<uint32_t,(~S0 & F1) | (~S1 & F0)>{});
|
||||
C<(~S0 & F1) | (~S1 & F0)>{});
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator|(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const&)
|
||||
operator|(MixedBits<S0,F0> const& m, C<S1>)
|
||||
{
|
||||
return make_mixed_bits(constant<uint32_t,S0 | S1>{},
|
||||
return make_mixed_bits(C<S0 | uint32_t(S1)>{},
|
||||
m.dynamic_int_,
|
||||
constant<uint32_t,~S1 & F0>{});
|
||||
C<F0 & ~uint32_t(S1)>{});
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator|(constant<TS1,S1> const& s, MixedBits<S0,D0,F0> const& m)
|
||||
operator|(C<S1> s, MixedBits<S0,F0> const& m)
|
||||
{
|
||||
return m | s;
|
||||
}
|
||||
|
||||
// Bitwise XOR
|
||||
template <uint32_t S0, class D0, uint32_t F0,
|
||||
uint32_t S1, class D1, uint32_t F1>
|
||||
template <uint32_t S0, uint32_t F0,
|
||||
uint32_t S1, uint32_t F1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator^(MixedBits<S0,D0,F0> const& m0, MixedBits<S1,D1,F1> const& m1)
|
||||
operator^(MixedBits<S0,F0> const& m0, MixedBits<S1,F1> const& m1)
|
||||
{
|
||||
// Truth table for (S0,D0,F0) ^ (S1,D1,F1) -> (S,D,F)
|
||||
// S0D0F0 | 0X0 | 001 | 011 | 1X0 |
|
||||
@@ -371,53 +370,53 @@ operator^(MixedBits<S0,D0,F0> const& m0, MixedBits<S1,D1,F1> const& m1)
|
||||
// 011 | 011 | 011 | 001 | 001 |
|
||||
// 1X0 | 1X0 | 011 | 001 | 0X0 |
|
||||
|
||||
return make_mixed_bits(constant<uint32_t,(~S0 & S1 & ~F0) | (S0 & ~S1 & ~F1)>{},
|
||||
return make_mixed_bits(C<(~S0 & S1 & ~F0) | (S0 & ~S1 & ~F1)>{},
|
||||
(S0 | m0.dynamic_int_) ^ (S1 | m1.dynamic_int_),
|
||||
constant<uint32_t,F0 | F1>{});
|
||||
C<F0 | F1>{});
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator^(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const&)
|
||||
operator^(MixedBits<S0,F0> const& m, C<S1>)
|
||||
{
|
||||
return make_mixed_bits(constant<uint32_t,(~S0 & S1 & ~F0) | (S0 & ~S1)>{},
|
||||
(S0 | m.dynamic_int_) ^ S1,
|
||||
constant<uint32_t,F0>{});
|
||||
return make_mixed_bits(C<(~S0 & uint32_t(S1) & ~F0) | (S0 & ~uint32_t(S1))>{},
|
||||
(S0 | m.dynamic_int_) ^ uint32_t(S1),
|
||||
C<F0>{});
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator^(constant<TS1,S1> const& s, MixedBits<S0,D0,F0> const& m)
|
||||
operator^(C<S1> s, MixedBits<S0,F0> const& m)
|
||||
{
|
||||
return m ^ s;
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator<<(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const&)
|
||||
operator<<(MixedBits<S0,F0> const& m, C<S1>)
|
||||
{
|
||||
return make_mixed_bits(constant<uint32_t,(S0 << S1)>{},
|
||||
return make_mixed_bits(C<(S0 << S1)>{},
|
||||
m.dynamic_int_ << S1,
|
||||
constant<uint32_t,(F0 << S1)>{});
|
||||
C<(F0 << S1)>{});
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
operator>>(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const&)
|
||||
operator>>(MixedBits<S0,F0> const& m, C<S1>)
|
||||
{
|
||||
return make_mixed_bits(constant<uint32_t,(S0 >> S1)>{},
|
||||
return make_mixed_bits(C<(S0 >> S1)>{},
|
||||
m.dynamic_int_ >> S1,
|
||||
constant<uint32_t,(F0 >> S1)>{});
|
||||
C<(F0 >> S1)>{});
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
shiftl(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const& s)
|
||||
shiftl(MixedBits<S0,F0> const& m, C<S1> s)
|
||||
{
|
||||
if constexpr (S1 >= 0) {
|
||||
return m << s;
|
||||
@@ -426,10 +425,10 @@ shiftl(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const& s)
|
||||
}
|
||||
}
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
shiftr(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const& s)
|
||||
shiftr(MixedBits<S0,F0> const& m, C<S1> s)
|
||||
{
|
||||
if constexpr (S1 >= 0) {
|
||||
return m >> s;
|
||||
@@ -442,24 +441,24 @@ shiftr(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const& s)
|
||||
// upcast and downcast
|
||||
//
|
||||
|
||||
template <uint32_t S0, class D0, uint32_t F0, class TS1, TS1 S1>
|
||||
template <uint32_t S0, uint32_t F0, auto S1>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
safe_div(MixedBits<S0,D0,F0> const& m, constant<TS1,S1> const& s)
|
||||
safe_div(MixedBits<S0,F0> const& m, C<S1> s)
|
||||
{
|
||||
static_assert(has_single_bit(S1), "Only divide MixedBits by powers of two.");
|
||||
return make_mixed_bits(safe_div(constant<uint32_t,S0>{}, s),
|
||||
static_assert(has_single_bit(uint32_t(S1)), "Only divide MixedBits by powers of two.");
|
||||
return make_mixed_bits(safe_div(C<S0>{}, s),
|
||||
safe_div(m.dynamic_int_, s),
|
||||
safe_div(constant<uint32_t,F0>{}, s));
|
||||
safe_div(C<F0>{}, s));
|
||||
}
|
||||
|
||||
template <uint32_t N, uint32_t S0, class D0, uint32_t F0>
|
||||
template <uint32_t N, uint32_t S0, uint32_t F0>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
upcast(MixedBits<S0,D0,F0> const& m)
|
||||
upcast(MixedBits<S0,F0> const& m)
|
||||
{
|
||||
static_assert(has_single_bit(N), "Only divide MixedBits by powers of two.");
|
||||
return safe_div(m, constant<uint32_t,N>{});
|
||||
return safe_div(m, C<N>{});
|
||||
}
|
||||
|
||||
template <uint32_t N, class T, __CUTE_REQUIRES(cute::is_integral<T>::value)>
|
||||
@@ -467,18 +466,18 @@ CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
upcast(T const& m)
|
||||
{
|
||||
return safe_div(m, constant<uint32_t,N>{});
|
||||
return safe_div(m, C<N>{});
|
||||
}
|
||||
|
||||
template <uint32_t N, uint32_t S0, class D0, uint32_t F0>
|
||||
template <uint32_t N, uint32_t S0, uint32_t F0>
|
||||
CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
downcast(MixedBits<S0,D0,F0> const& m)
|
||||
downcast(MixedBits<S0,F0> const& m)
|
||||
{
|
||||
static_assert(has_single_bit(N), "Only scale MixedBits by powers of two.");
|
||||
return make_mixed_bits(constant<uint32_t,S0 * N>{},
|
||||
return make_mixed_bits(C<S0 * N>{},
|
||||
m.dynamic_int_ * N,
|
||||
constant<uint32_t,F0 * N>{});
|
||||
C<F0 * N>{});
|
||||
}
|
||||
|
||||
template <uint32_t N, class T, __CUTE_REQUIRES(cute::is_integral<T>::value)>
|
||||
@@ -486,7 +485,7 @@ CUTE_HOST_DEVICE constexpr
|
||||
auto
|
||||
downcast(T const& m)
|
||||
{
|
||||
return m * constant<uint32_t, N>{};
|
||||
return m * C<N>{};
|
||||
}
|
||||
|
||||
//
|
||||
@@ -525,17 +524,17 @@ to_mixed_bits(Layout const& layout, Coord const& coord)
|
||||
// Display utilities
|
||||
//
|
||||
|
||||
template <uint32_t S, class D, uint32_t F>
|
||||
CUTE_HOST_DEVICE void print(MixedBits<S,D,F> const& m)
|
||||
template <uint32_t S, uint32_t F>
|
||||
CUTE_HOST_DEVICE void print(MixedBits<S,F> const& m)
|
||||
{
|
||||
printf("M_%u|(%u&%u)=%u", S, uint32_t(m.dynamic_int_), F, uint32_t(m));
|
||||
printf("M_%u|(%u&%u)=%u", S, m.dynamic_int_, F, uint32_t(m));
|
||||
}
|
||||
|
||||
#if !defined(__CUDACC_RTC__)
|
||||
template <uint32_t S, class D, uint32_t F>
|
||||
CUTE_HOST std::ostream& operator<<(std::ostream& os, MixedBits<S,D,F> const& m)
|
||||
CUTE_HOST std::ostream& operator<<(std::ostream& os, MixedBits<S,F> const& m)
|
||||
{
|
||||
return os << "M_" << S << "|(" << uint32_t(m.dynamic_int_) << "&" << F << ")=" << uint32_t(m);
|
||||
return os << "M_" << S << "|(" << m.dynamic_int_ << "&" << F << ")=" << uint32_t(m);
|
||||
}
|
||||
|
||||
template <int B, int M, int S>
|
||||
|
||||
Reference in New Issue
Block a user