Block evaluation for TensorGenerator + TensorReverse + fixed bug in tensor reverse op

This commit is contained in:
Eugene Zhulenev
2019-10-10 10:56:58 -07:00
parent b03eb63d7c
commit a411e9f344
5 changed files with 303 additions and 40 deletions

View File

@@ -369,6 +369,48 @@ static void test_eval_tensor_chipping() {
[&chipped_dims]() { return RandomBlock<Layout>(chipped_dims, 1, 10); });
}
template <typename T, int NumDims, int Layout>
static void test_eval_tensor_generator() {
DSizes<Index, NumDims> dims = RandomDims<NumDims>(10, 20);
Tensor<T, NumDims, Layout> input(dims);
input.setRandom();
auto generator = [](const array<Index, NumDims>& dims) -> T {
T result = static_cast<T>(0);
for (int i = 0; i < NumDims; ++i) {
result += static_cast<T>((i + 1) * dims[i]);
}
return result;
};
VerifyBlockEvaluator<T, NumDims, Layout>(
input.generate(generator),
[&dims]() { return FixedSizeBlock(dims); });
VerifyBlockEvaluator<T, NumDims, Layout>(
input.generate(generator),
[&dims]() { return RandomBlock<Layout>(dims, 1, 10); });
}
template <typename T, int NumDims, int Layout>
static void test_eval_tensor_reverse() {
DSizes<Index, NumDims> dims = RandomDims<NumDims>(10, 20);
Tensor<T, NumDims, Layout> input(dims);
input.setRandom();
// Randomly reverse dimensions.
Eigen::DSizes<bool, NumDims> reverse;
for (int i = 0; i < NumDims; ++i) reverse[i] = internal::random<bool>();
VerifyBlockEvaluator<T, NumDims, Layout>(
input.reverse(reverse),
[&dims]() { return FixedSizeBlock(dims); });
VerifyBlockEvaluator<T, NumDims, Layout>(
input.reverse(reverse),
[&dims]() { return RandomBlock<Layout>(dims, 1, 10); });
}
template <typename T, int Layout>
static void test_eval_tensor_reshape_with_bcast() {
Index dim = internal::random<Index>(1, 100);
@@ -573,6 +615,8 @@ EIGEN_DECLARE_TEST(cxx11_tensor_block_eval) {
CALL_SUBTESTS_DIMS_LAYOUTS(test_eval_tensor_select);
CALL_SUBTESTS_DIMS_LAYOUTS(test_eval_tensor_padding);
CALL_SUBTESTS_DIMS_LAYOUTS(test_eval_tensor_chipping);
CALL_SUBTESTS_DIMS_LAYOUTS(test_eval_tensor_generator);
CALL_SUBTESTS_DIMS_LAYOUTS(test_eval_tensor_reverse);
CALL_SUBTESTS_LAYOUTS(test_eval_tensor_reshape_with_bcast);
CALL_SUBTESTS_LAYOUTS(test_eval_tensor_forced_eval);

View File

@@ -539,7 +539,7 @@ static void test_execute_reverse_rvalue(Device d)
// Reverse half of the dimensions.
Eigen::array<bool, NumDims> reverse;
for (int i = 0; i < NumDims; ++i) reverse[i] = (dims[i] % 2 == 0);
for (int i = 0; i < NumDims; ++i) reverse[i] = internal::random<bool>();
const auto expr = src.reverse(reverse);
@@ -756,16 +756,16 @@ EIGEN_DECLARE_TEST(cxx11_tensor_executor) {
CALL_SUBTEST_COMBINATIONS_V2(12, test_execute_broadcasting_of_forced_eval, float, 4);
CALL_SUBTEST_COMBINATIONS_V2(12, test_execute_broadcasting_of_forced_eval, float, 5);
CALL_SUBTEST_COMBINATIONS_V1(13, test_execute_generator_op, float, 2);
CALL_SUBTEST_COMBINATIONS_V1(13, test_execute_generator_op, float, 3);
CALL_SUBTEST_COMBINATIONS_V1(13, test_execute_generator_op, float, 4);
CALL_SUBTEST_COMBINATIONS_V1(13, test_execute_generator_op, float, 5);
CALL_SUBTEST_COMBINATIONS_V2(13, test_execute_generator_op, float, 2);
CALL_SUBTEST_COMBINATIONS_V2(13, test_execute_generator_op, float, 3);
CALL_SUBTEST_COMBINATIONS_V2(13, test_execute_generator_op, float, 4);
CALL_SUBTEST_COMBINATIONS_V2(13, test_execute_generator_op, float, 5);
CALL_SUBTEST_COMBINATIONS_V1(14, test_execute_reverse_rvalue, float, 1);
CALL_SUBTEST_COMBINATIONS_V1(14, test_execute_reverse_rvalue, float, 2);
CALL_SUBTEST_COMBINATIONS_V1(14, test_execute_reverse_rvalue, float, 3);
CALL_SUBTEST_COMBINATIONS_V1(14, test_execute_reverse_rvalue, float, 4);
CALL_SUBTEST_COMBINATIONS_V1(14, test_execute_reverse_rvalue, float, 5);
CALL_SUBTEST_COMBINATIONS_V2(14, test_execute_reverse_rvalue, float, 1);
CALL_SUBTEST_COMBINATIONS_V2(14, test_execute_reverse_rvalue, float, 2);
CALL_SUBTEST_COMBINATIONS_V2(14, test_execute_reverse_rvalue, float, 3);
CALL_SUBTEST_COMBINATIONS_V2(14, test_execute_reverse_rvalue, float, 4);
CALL_SUBTEST_COMBINATIONS_V2(14, test_execute_reverse_rvalue, float, 5);
CALL_ASYNC_SUBTEST_COMBINATIONS(15, test_async_execute_unary_expr, float, 3);
CALL_ASYNC_SUBTEST_COMBINATIONS(15, test_async_execute_unary_expr, float, 4);