Updates for CUTLASS 3.5.0 (#1468)
This commit is contained in:
@@ -35,22 +35,22 @@
|
||||
|
||||
#include <cute/tensor.hpp>
|
||||
|
||||
template <class Layout, class CoSizeHi>
|
||||
template <class Layout, class CoTarget>
|
||||
void
|
||||
test_complement(Layout const& layout, CoSizeHi const& cosize_hi)
|
||||
test_complement(Layout const& layout, CoTarget const& cotarget)
|
||||
{
|
||||
using namespace cute;
|
||||
|
||||
auto result = complement(layout, cosize_hi);
|
||||
auto result = complement(layout, cotarget);
|
||||
|
||||
CUTLASS_TRACE_HOST("complement(" << layout << ", " << cosize_hi << ") => " << result);
|
||||
CUTLASS_TRACE_HOST("complement(" << layout << ", " << cotarget << ") => " << result);
|
||||
|
||||
auto completed = make_layout(layout, result);
|
||||
|
||||
// Lower-bound on the codomain size of the layout ++ complement (1)
|
||||
EXPECT_GE(cosize(completed), cosize_hi);
|
||||
EXPECT_GE(cosize(completed), size(cotarget));
|
||||
// Upper-bound on the codomain size of the complement (2)
|
||||
EXPECT_LE(cosize(result), cute::round_up(cosize_hi, cosize(layout)));
|
||||
EXPECT_LE(cosize(result), cute::round_up(size(cotarget), cosize(layout)));
|
||||
|
||||
// Post-condition on the codomain of the complement
|
||||
for (int i = 1; i < size(result); ++i) {
|
||||
@@ -62,9 +62,9 @@ test_complement(Layout const& layout, CoSizeHi const& cosize_hi)
|
||||
|
||||
// Other observations
|
||||
EXPECT_LE(size(result), cosize(result)); // As a result of the ordered condition (3)
|
||||
EXPECT_GE(size(result), cosize_hi / size(filter(layout)));
|
||||
EXPECT_GE(size(result), size(cotarget) / size(filter(layout)));
|
||||
EXPECT_LE(cosize(completed), cosize(result) + cosize(layout));
|
||||
EXPECT_GE(cosize(result), cosize_hi / size(filter(layout)));
|
||||
EXPECT_GE(cosize(result), size(cotarget) / 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
|
||||
}
|
||||
@@ -90,6 +90,8 @@ TEST(CuTe_core, Complement)
|
||||
|
||||
test_complement(layout);
|
||||
test_complement(layout, Int<2>{});
|
||||
test_complement(layout, Int<5>{});
|
||||
test_complement(layout, make_shape(Int<2>{}, 2));
|
||||
}
|
||||
|
||||
{
|
||||
@@ -97,6 +99,8 @@ TEST(CuTe_core, Complement)
|
||||
|
||||
test_complement(layout);
|
||||
test_complement(layout, Int<2>{});
|
||||
test_complement(layout, Int<5>{});
|
||||
test_complement(layout, make_shape(Int<2>{}, 2));
|
||||
}
|
||||
|
||||
{
|
||||
@@ -105,6 +109,8 @@ TEST(CuTe_core, Complement)
|
||||
test_complement(layout, Int<1>{});
|
||||
test_complement(layout, Int<2>{});
|
||||
test_complement(layout, Int<8>{});
|
||||
test_complement(layout, Int<5>{});
|
||||
test_complement(layout, make_shape(Int<2>{}, 2));
|
||||
}
|
||||
|
||||
{
|
||||
@@ -130,6 +136,7 @@ TEST(CuTe_core, Complement)
|
||||
test_complement(layout);
|
||||
test_complement(layout, Int<16>{});
|
||||
test_complement(layout, Int<19>{});
|
||||
test_complement(layout, make_shape(Int<2>{}, 2));
|
||||
}
|
||||
|
||||
{
|
||||
@@ -138,6 +145,7 @@ TEST(CuTe_core, Complement)
|
||||
test_complement(layout, Int<1>{});
|
||||
test_complement(layout);
|
||||
test_complement(layout, Int<17>{});
|
||||
test_complement(layout, make_shape(Int<2>{}, 2));
|
||||
}
|
||||
|
||||
{
|
||||
@@ -193,8 +201,8 @@ TEST(CuTe_core, Complement)
|
||||
|
||||
// Fails due to non-injective layout
|
||||
// {
|
||||
// auto layout = make_layout(Shape<Shape<_2,_2>,Shape<_2, _2>>{},
|
||||
// Stride<Stride<_1,_8>,Stride<_8,_4>>{});
|
||||
// auto layout = make_layout(Shape <Shape <_2,_2>,Shape <_2,_2>>{},
|
||||
// Stride<Stride<_1,_8>,Stride<_8,_4>>{});
|
||||
|
||||
// test_complement(layout);
|
||||
// }
|
||||
@@ -289,4 +297,11 @@ TEST(CuTe_core, Complement)
|
||||
|
||||
test_complement(layout);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = make_layout(Int<64>{});
|
||||
|
||||
test_complement(layout, make_shape(Int<32>{}, Int<4>{}, Int<4>{}));
|
||||
test_complement(layout, make_shape(Int<32>{}, Int<4>{}, 4));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -212,13 +212,12 @@ TEST(CuTe_core, Composition)
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
// FAILS due to b not "dividing into" a properly
|
||||
//{
|
||||
// auto a = make_layout(Shape<_4,_3>{});
|
||||
// auto b = make_layout(Shape<_6>{});
|
||||
{
|
||||
auto a = make_layout(Shape<_4,_3>{});
|
||||
auto b = make_layout(Shape<_6>{});
|
||||
|
||||
// test_composition(a, b);
|
||||
//}
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
{
|
||||
auto a = make_layout(Shape<_4,_3>{});
|
||||
@@ -234,13 +233,12 @@ TEST(CuTe_core, Composition)
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
// FAILS due to b not "dividing into" a properly
|
||||
//{
|
||||
// auto a = make_layout(Shape<_4,_3>{});
|
||||
// auto b = make_layout(Shape<_4,_3>{}, Stride<_3,_1>{});
|
||||
{
|
||||
auto a = make_layout(Shape<_4,_3>{});
|
||||
auto b = make_layout(Shape<_4,_3>{}, Stride<_3,_1>{});
|
||||
|
||||
// test_composition(a, b);
|
||||
//}
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
{
|
||||
auto a = make_layout(Shape<_4,_3>{}, Stride<_3,_1>{});
|
||||
@@ -523,4 +521,21 @@ TEST(CuTe_core, Composition)
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
CUTLASS_TRACE_HOST("BETA: Tuple strides" );
|
||||
CUTLASS_TRACE_HOST("-------------------------------");
|
||||
|
||||
{
|
||||
auto a = make_layout(Shape<_4,_4>{}, Stride<_4,_1>{});
|
||||
auto b = make_layout(Shape<_4,_4>{}, Stride<E<1>,E<0>>{});
|
||||
|
||||
test_composition(a, b);
|
||||
}
|
||||
|
||||
{
|
||||
auto a = make_layout(Shape<_4,Shape<_2,_3>>{}, Stride<_6,Stride<_3,_1>>{});
|
||||
auto b = make_layout(Shape<_2,_4>{}, Stride<E<1,1>,E<0>>{});
|
||||
|
||||
test_composition(a, b);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -227,27 +227,42 @@ TEST(CuTe_core, Logical_divide)
|
||||
ASSERT_TRUE(decltype(stride<1>(result) == Int<48>{})::value);
|
||||
}
|
||||
|
||||
// DISALLOWED
|
||||
//{
|
||||
//auto layout = make_layout(make_shape(128,4,3), make_stride(1,512,0));
|
||||
//auto tile = Layout<_32>{};
|
||||
{
|
||||
auto layout = make_layout(make_shape(Int<32>{}, Int<4>{}, 4));
|
||||
auto tile = Layout<_64>{};
|
||||
|
||||
//test_logical_divide(layout, tile);
|
||||
//}
|
||||
test_logical_divide(layout, tile);
|
||||
|
||||
//{
|
||||
//auto layout = make_layout(make_shape(128,4,3), make_stride(1,512,0));
|
||||
//auto tile = Layout<_32,_2>{};
|
||||
// Enforcement of result
|
||||
auto result = logical_divide(layout, tile);
|
||||
ASSERT_TRUE(bool( shape(result) == make_shape (_64{}, make_shape ( _2{}, 4))));
|
||||
ASSERT_TRUE(bool(stride(result) == make_stride( _1{}, make_stride(_64{},_128{}))));
|
||||
}
|
||||
|
||||
//CUTLASS_TRACE_HOST("complement: " << complement(tile, size(layout)));
|
||||
//test_logical_divide(layout, tile);
|
||||
//}
|
||||
|
||||
//{
|
||||
//auto layout = make_layout(make_shape(16,4,3), make_stride(1,512,0));
|
||||
//auto tile = Layout<_32>{};
|
||||
//
|
||||
// ALLOWED, but dangerous due to the dynamic lhs shapes
|
||||
// Consider disallowing...
|
||||
//
|
||||
|
||||
//CUTLASS_TRACE_HOST("complement: " << complement(tile, size(layout)));
|
||||
//test_logical_divide(layout, tile);
|
||||
//}
|
||||
{
|
||||
auto layout = make_layout(make_shape(128,4,3), make_stride(1,512,0));
|
||||
auto tile = Layout<_32>{};
|
||||
|
||||
test_logical_divide(layout, tile);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = make_layout(make_shape(128,4,3), make_stride(1,512,0));
|
||||
auto tile = Layout<_32,_2>{};
|
||||
|
||||
test_logical_divide(layout, tile);
|
||||
}
|
||||
|
||||
{
|
||||
auto layout = make_layout(make_shape(16,4,3), make_stride(1,512,0));
|
||||
auto tile = Layout<_32>{};
|
||||
|
||||
test_logical_divide(layout, tile);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user