1 // -*- C++ -*- 2 //===-- nth_element.pass.cpp ----------------------------------------------===// 3 // 4 // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions. 5 // See https://llvm.org/LICENSE.txt for license information. 6 // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception 7 // 8 //===----------------------------------------------------------------------===// 9 10 #include "support/pstl_test_config.h" 11 12 #ifdef PSTL_STANDALONE_TESTS 13 #include <algorithm> 14 #include <iostream> 15 #include "pstl/execution" 16 #include "pstl/algorithm" 17 18 #else 19 #include <execution> 20 #include <algorithm> 21 #endif // PSTL_STANDALONE_TESTS 22 23 #include "support/utils.h" 24 25 using namespace TestUtils; 26 27 // User defined type with minimal requirements 28 template <typename T> 29 struct DataType 30 { 31 explicit DataType(int32_t k) : my_val(k) {} 32 DataType(DataType&& input) 33 { 34 my_val = std::move(input.my_val); 35 input.my_val = T(0); 36 } 37 DataType& 38 operator=(DataType&& input) 39 { 40 my_val = std::move(input.my_val); 41 input.my_val = T(0); 42 return *this; 43 } 44 T 45 get_val() const 46 { 47 return my_val; 48 } 49 50 friend std::ostream& 51 operator<<(std::ostream& stream, const DataType<T>& input) 52 { 53 return stream << input.my_val; 54 } 55 56 private: 57 T my_val; 58 }; 59 60 template <typename T> 61 bool 62 is_equal(const DataType<T>& x, const DataType<T>& y) 63 { 64 return x.get_val() == y.get_val(); 65 } 66 67 template <typename T> 68 bool 69 is_equal(const T& x, const T& y) 70 { 71 return x == y; 72 } 73 74 struct test_one_policy 75 { 76 #if _PSTL_ICC_17_VC141_TEST_SIMD_LAMBDA_DEBUG_32_BROKEN || \ 77 _PSTL_ICC_16_VC14_TEST_SIMD_LAMBDA_DEBUG_32_BROKEN // dummy specialization by policy type, in case of broken configuration 78 template <typename Iterator1, typename Size, typename Generator1, typename Generator2, typename Compare> 79 typename std::enable_if<is_same_iterator_category<Iterator1, std::random_access_iterator_tag>::value, void>::type 80 operator()(pstl::execution::unsequenced_policy, Iterator1 first1, Iterator1 last1, Iterator1 first2, 81 Iterator1 last2, Size n, Size m, Generator1 generator1, Generator2 generator2, Compare comp) 82 { 83 } 84 template <typename Iterator1, typename Size, typename Generator1, typename Generator2, typename Compare> 85 typename std::enable_if<is_same_iterator_category<Iterator1, std::random_access_iterator_tag>::value, void>::type 86 operator()(pstl::execution::parallel_unsequenced_policy, Iterator1 first1, Iterator1 last1, Iterator1 first2, 87 Iterator1 last2, Size n, Size m, Generator1 generator1, Generator2 generator2, Compare comp) 88 { 89 } 90 #endif 91 92 // nth_element works only with random access iterators 93 template <typename Policy, typename Iterator1, typename Size, typename Generator1, typename Generator2, 94 typename Compare> 95 typename std::enable_if<is_same_iterator_category<Iterator1, std::random_access_iterator_tag>::value, void>::type 96 operator()(Policy&& exec, Iterator1 first1, Iterator1 last1, Iterator1 first2, Iterator1 last2, Size n, Size m, 97 Generator1 generator1, Generator2 generator2, Compare comp) 98 { 99 100 using T = typename std::iterator_traits<Iterator1>::value_type; 101 const Iterator1 mid1 = std::next(first1, m); 102 const Iterator1 mid2 = std::next(first2, m); 103 104 fill_data(first1, mid1, generator1); 105 fill_data(mid1, last1, generator2); 106 fill_data(first2, mid2, generator1); 107 fill_data(mid2, last2, generator2); 108 std::nth_element(first1, mid1, last1, comp); 109 std::nth_element(exec, first2, mid2, last2, comp); 110 if (m > 0 && m < n) 111 { 112 EXPECT_TRUE(is_equal(*mid1, *mid2), "wrong result from nth_element with predicate"); 113 } 114 EXPECT_TRUE(std::find_first_of(first2, mid2, mid2, last2, [comp](T& x, T& y) { return comp(y, x); }) == mid2, 115 "wrong effect from nth_element with predicate"); 116 } 117 118 template <typename Policy, typename Iterator1, typename Size, typename Generator1, typename Generator2, 119 typename Compare> 120 typename std::enable_if<!is_same_iterator_category<Iterator1, std::random_access_iterator_tag>::value, void>::type 121 operator()(Policy&& exec, Iterator1 first1, Iterator1 last1, Iterator1 first2, Iterator1 last2, Size n, Size m, 122 Generator1 generator1, Generator2 generator2, Compare comp) 123 { 124 } 125 }; 126 127 template <typename T, typename Generator1, typename Generator2, typename Compare> 128 void 129 test_by_type(Generator1 generator1, Generator2 generator2, Compare comp) 130 { 131 using namespace std; 132 size_t max_size = 10000; 133 Sequence<T> in1(max_size, [](size_t v) { return T(v); }); 134 Sequence<T> exp(max_size, [](size_t v) { return T(v); }); 135 size_t m; 136 137 for (size_t n = 0; n <= max_size; n = n <= 16 ? n + 1 : size_t(3.1415 * n)) 138 { 139 m = 0; 140 invoke_on_all_policies(test_one_policy(), exp.begin(), exp.begin() + n, in1.begin(), in1.begin() + n, n, m, 141 generator1, generator2, comp); 142 m = n / 7; 143 invoke_on_all_policies(test_one_policy(), exp.begin(), exp.begin() + n, in1.begin(), in1.begin() + n, n, m, 144 generator1, generator2, comp); 145 m = 3 * n / 5; 146 invoke_on_all_policies(test_one_policy(), exp.begin(), exp.begin() + n, in1.begin(), in1.begin() + n, n, m, 147 generator1, generator2, comp); 148 } 149 invoke_on_all_policies(test_one_policy(), exp.begin(), exp.begin() + max_size, in1.begin(), in1.begin() + max_size, 150 max_size, max_size, generator1, generator2, comp); 151 } 152 153 template <typename T> 154 struct test_non_const 155 { 156 template <typename Policy, typename Iterator> 157 void 158 operator()(Policy&& exec, Iterator iter) 159 { 160 invoke_if(exec, [&]() { nth_element(exec, iter, iter, iter, non_const(std::less<T>())); }); 161 } 162 }; 163 164 int32_t 165 main() 166 { 167 test_by_type<int32_t>([](int32_t i) { return 10 * i; }, [](int32_t i) { return i + 1; }, std::less<int32_t>()); 168 test_by_type<int32_t>([](int32_t) { return 0; }, [](int32_t) { return 0; }, std::less<int32_t>()); 169 170 test_by_type<float64_t>([](int32_t i) { return -2 * i; }, [](int32_t i) { return -(2 * i + 1); }, 171 [](const float64_t x, const float64_t y) { return x > y; }); 172 173 test_by_type<DataType<float32_t>>( 174 [](int32_t i) { return DataType<float32_t>(2 * i + 1); }, [](int32_t i) { return DataType<float32_t>(2 * i); }, 175 [](const DataType<float32_t>& x, const DataType<float32_t>& y) { return x.get_val() < y.get_val(); }); 176 177 test_algo_basic_single<int32_t>(run_for_rnd<test_non_const<int32_t>>()); 178 179 std::cout << done() << std::endl; 180 return 0; 181 } 182