Updates for 3.4 release. (#1305)
This commit is contained in:
@@ -38,9 +38,10 @@
|
||||
#include <vector>
|
||||
#include <numeric>
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
#include <cute/container/bit_field.hpp>
|
||||
|
||||
#include <cute/algorithm/tuple_algorithms.hpp>
|
||||
|
||||
using namespace cute;
|
||||
|
||||
TEST(CuTe_core, Bitfield)
|
||||
|
||||
@@ -43,26 +43,30 @@ test_complement(Layout const& layout, CoSizeHi const& cosize_hi)
|
||||
|
||||
auto result = complement(layout, cosize_hi);
|
||||
|
||||
CUTLASS_TRACE_HOST("complement( " << layout << ", " << cosize_hi << ") => " << result);
|
||||
CUTLASS_TRACE_HOST("complement(" << layout << ", " << cosize_hi << ") => " << result);
|
||||
|
||||
// Post-condition on the domain size of the complement (1)
|
||||
EXPECT_GE( size(result), cosize_hi / size(filter(layout)));
|
||||
// Post-condition on the codomain size of the complement (2)
|
||||
EXPECT_LE(cosize(result), cute::ceil_div(cosize_hi, cosize(layout)) * cosize(layout));
|
||||
auto completed = make_layout(layout, result);
|
||||
|
||||
// Lower-bound on the codomain size of the layout ++ complement (1)
|
||||
EXPECT_GE(cosize(completed), cosize_hi);
|
||||
// Upper-bound on the codomain size of the complement (2)
|
||||
EXPECT_LE(cosize(result), cute::round_up(cosize_hi, cosize(layout)));
|
||||
|
||||
// Post-condition on the codomain of the complement
|
||||
for (int i = 1; i < size(result); ++i) {
|
||||
EXPECT_LT(result(i-1), result(i)); // Ordered (3)
|
||||
for (int j = 0; j < size(layout); ++j) {
|
||||
EXPECT_NE(result(i), layout(j)); // Complemented (4)
|
||||
EXPECT_NE(result(i), layout(j)); // Disjoint (4)
|
||||
}
|
||||
}
|
||||
|
||||
// Other observations
|
||||
EXPECT_LE(size(result),cosize(result)); // As a result of the ordered condition (3)
|
||||
EXPECT_GE(cosize(result), cosize_hi / size(filter(layout))); // As a result of (1) (2) and (5)
|
||||
if constexpr (is_static<decltype(stride(make_layout(layout,result)))>::value) { // If we can apply complement again
|
||||
EXPECT_EQ(size(complement(make_layout(layout,result))), 1); // There's no more codomain left over
|
||||
EXPECT_LE(size(result), cosize(result)); // As a result of the ordered condition (3)
|
||||
EXPECT_GE(size(result), cosize_hi / size(filter(layout)));
|
||||
EXPECT_LE(cosize(completed), cosize(result) + cosize(layout));
|
||||
EXPECT_GE(cosize(result), cosize_hi / size(filter(layout)));
|
||||
if constexpr (is_static<decltype(stride(completed))>::value) { // If we can apply complement again
|
||||
EXPECT_EQ(size(complement(completed)), 1); // There's no more codomain left over
|
||||
}
|
||||
}
|
||||
|
||||
@@ -125,6 +129,7 @@ TEST(CuTe_core, Complement)
|
||||
test_complement(layout, Int<1>{});
|
||||
test_complement(layout);
|
||||
test_complement(layout, Int<16>{});
|
||||
test_complement(layout, Int<19>{});
|
||||
}
|
||||
|
||||
{
|
||||
@@ -153,6 +158,12 @@ TEST(CuTe_core, Complement)
|
||||
test_complement(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_2,_4>, Stride<_1,_6>>{};
|
||||
|
||||
test_complement(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_2,_4,_8>, Stride<_8,_1,_64>>{};
|
||||
|
||||
@@ -167,26 +178,34 @@ TEST(CuTe_core, Complement)
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = make_layout(Shape<Shape<_2,_2>,Shape<_2, _2>>{},
|
||||
auto layout = make_layout(Shape <Shape <_2,_2>,Shape <_2, _2>>{},
|
||||
Stride<Stride<_1,_4>,Stride<_8,_32>>{});
|
||||
|
||||
test_complement(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = make_layout(Shape<Shape<_2,_2>,Shape<_2, _2>>{},
|
||||
auto layout = make_layout(Shape <Shape <_2, _2>,Shape <_2,_2>>{},
|
||||
Stride<Stride<_1,_32>,Stride<_8,_4>>{});
|
||||
|
||||
test_complement(layout);
|
||||
}
|
||||
|
||||
// Fails due to non-injective input
|
||||
//{
|
||||
//auto layout = make_layout(Shape<Shape<_2,_2>,Shape<_2, _2>>{},
|
||||
// Fails due to non-injective layout
|
||||
// {
|
||||
// auto layout = make_layout(Shape<Shape<_2,_2>,Shape<_2, _2>>{},
|
||||
// Stride<Stride<_1,_8>,Stride<_8,_4>>{});
|
||||
|
||||
//test_complement(layout);
|
||||
//}
|
||||
// test_complement(layout);
|
||||
// }
|
||||
|
||||
// Fails due to non-injective layout
|
||||
// {
|
||||
// auto layout = Layout<Shape<_2,_2>, Stride<_2,_3>>{};
|
||||
|
||||
// test_complement(layout);
|
||||
// test_complement(layout, Int<19>{});
|
||||
// }
|
||||
|
||||
{
|
||||
auto layout = Layout<Shape<_4,_6>, Stride<_1,_6>>{};
|
||||
|
||||
@@ -42,8 +42,8 @@ using namespace cute;
|
||||
|
||||
template <class LayoutA, class LayoutB>
|
||||
void
|
||||
test_composition(const LayoutA& layoutA,
|
||||
const LayoutB& layoutB)
|
||||
test_composition(LayoutA const& layoutA,
|
||||
LayoutB const& layoutB)
|
||||
{
|
||||
auto layoutR = composition(layoutA, layoutB);
|
||||
|
||||
@@ -52,14 +52,12 @@ test_composition(const LayoutA& layoutA,
|
||||
CUTLASS_TRACE_HOST(" => ");
|
||||
CUTLASS_TRACE_HOST(layoutR);
|
||||
|
||||
// Test that layout R is compatible with layout B
|
||||
// Test that layout B is compatible with layout R
|
||||
EXPECT_TRUE(compatible(layoutB, layoutR));
|
||||
|
||||
// True post-condition: Every coordinate c of layoutB with L1D(c) < size(layoutR) is a coordinate of layoutR.
|
||||
|
||||
// Test that R(c) = A(B(c)) for all coordinates c in layoutR
|
||||
for (int i = 0; i < size(layoutR); ++i) {
|
||||
EXPECT_EQ(layoutR(i), layoutA(layoutB(i)));
|
||||
// Test that R(c) = A(B(c)) for all coordinates c in layoutB
|
||||
for (int c = 0; c < size(layoutB); ++c) {
|
||||
EXPECT_EQ(layoutR(c), layoutA(layoutB(c)));
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -45,10 +45,10 @@ test_logical_divide(LayoutA const& layoutA,
|
||||
auto layoutR = logical_divide(layoutA, layoutB);
|
||||
|
||||
CUTLASS_TRACE_HOST("test_logical_divide()");
|
||||
CUTLASS_TRACE_HOST(shape(layoutA) << " / " << shape(layoutB) << " => " << shape(layoutR) );
|
||||
CUTLASS_TRACE_HOST( shape(layoutA) << " / " << shape(layoutB) << " => " << shape(layoutR));
|
||||
CUTLASS_TRACE_HOST(stride(layoutA) << " " << stride(layoutB) << " => " << stride(layoutR));
|
||||
|
||||
// Test that layout R is compatible with layout B
|
||||
// Test that layout B is compatible with layout R_0
|
||||
ASSERT_EQ(rank(layoutR), 2);
|
||||
ASSERT_TRUE(compatible(layoutB, layout<0>(layoutR)));
|
||||
}
|
||||
@@ -186,10 +186,10 @@ TEST(CuTe_core, Logical_divide)
|
||||
|
||||
// Enforcement for dynamic cases
|
||||
auto result = logical_divide(layout, tile);
|
||||
static_assert(decltype(shape<0>(result) == Int<32>{})::value);
|
||||
static_assert(decltype(stride<0>(result) == Int<1>{})::value);
|
||||
assert(shape<1>(result) == 1);
|
||||
static_assert(decltype(stride<1>(result) == Int<32>{})::value);
|
||||
ASSERT_TRUE(decltype(shape<0>(result) == Int<32>{})::value);
|
||||
ASSERT_TRUE(decltype(stride<0>(result) == Int<1>{})::value);
|
||||
ASSERT_TRUE(shape<1>(result) == 1);
|
||||
ASSERT_TRUE(decltype(stride<1>(result) == Int<32>{})::value);
|
||||
}
|
||||
|
||||
{
|
||||
@@ -200,10 +200,10 @@ TEST(CuTe_core, Logical_divide)
|
||||
|
||||
// Enforcement for dynamic cases
|
||||
auto result = logical_divide(layout, tile);
|
||||
static_assert(decltype(shape<0>(result) == Int<32>{})::value);
|
||||
static_assert(decltype(stride<0>(result) == Int<1>{})::value);
|
||||
assert(shape<1>(result) == 2);
|
||||
static_assert(decltype(stride<1>(result) == Int<32>{})::value);
|
||||
ASSERT_TRUE(decltype(shape<0>(result) == Int<32>{})::value);
|
||||
ASSERT_TRUE(decltype(stride<0>(result) == Int<1>{})::value);
|
||||
ASSERT_TRUE(shape<1>(result) == 2);
|
||||
ASSERT_TRUE(decltype(stride<1>(result) == Int<32>{})::value);
|
||||
}
|
||||
|
||||
{
|
||||
@@ -221,10 +221,10 @@ TEST(CuTe_core, Logical_divide)
|
||||
|
||||
// Enforcement for dynamic cases
|
||||
auto result = logical_divide(layout, tile);
|
||||
static_assert(decltype(shape<0>(result) == Int<48>{})::value);
|
||||
static_assert(decltype(stride<0>(result) == Int<1>{})::value);
|
||||
assert(shape<1>(result) == 1);
|
||||
static_assert(decltype(stride<1>(result) == Int<48>{})::value);
|
||||
ASSERT_TRUE(decltype(shape<0>(result) == Int<48>{})::value);
|
||||
ASSERT_TRUE(decltype(stride<0>(result) == Int<1>{})::value);
|
||||
ASSERT_TRUE(shape<1>(result) == 1);
|
||||
ASSERT_TRUE(decltype(stride<1>(result) == Int<48>{})::value);
|
||||
}
|
||||
|
||||
// DISALLOWED
|
||||
|
||||
@@ -46,13 +46,9 @@ test_logical_product(LayoutA const& layoutA,
|
||||
CUTLASS_TRACE_HOST(shape(layoutA) << " x " << shape(layoutB) << " => " << shape(layoutR) );
|
||||
CUTLASS_TRACE_HOST(stride(layoutA) << " " << stride(layoutB) << " => " << stride(layoutR));
|
||||
|
||||
// Test that layout R is compatible with layout B
|
||||
ASSERT_EQ(rank(layoutR), 2);
|
||||
//assert(compatible(layoutB, layout<0>(layoutR)));
|
||||
//assert(consistent(layoutA, layout<1>(layoutR)));
|
||||
|
||||
// True post-condition:
|
||||
|
||||
ASSERT_TRUE(layoutA == layout<0>(layoutR));
|
||||
ASSERT_TRUE(compatible(layoutB, layout<1>(layoutR)));
|
||||
}
|
||||
|
||||
TEST(CuTe_core, Logical_product)
|
||||
|
||||
Reference in New Issue
Block a user