diff --git a/benchmarks/CMakeLists.txt b/benchmarks/CMakeLists.txt index 2903ee13245..8a463a154f8 100644 --- a/benchmarks/CMakeLists.txt +++ b/benchmarks/CMakeLists.txt @@ -110,6 +110,7 @@ endfunction() add_benchmark(bitset_to_string src/bitset_to_string.cpp) add_benchmark(find_and_count src/find_and_count.cpp) +add_benchmark(find_first_of src/find_first_of.cpp) add_benchmark(locale_classic src/locale_classic.cpp) add_benchmark(minmax_element src/minmax_element.cpp) add_benchmark(path_lexically_normal src/path_lexically_normal.cpp) diff --git a/benchmarks/src/find_first_of.cpp b/benchmarks/src/find_first_of.cpp new file mode 100644 index 00000000000..170793f1e58 --- /dev/null +++ b/benchmarks/src/find_first_of.cpp @@ -0,0 +1,46 @@ +// Copyright (c) Microsoft Corporation. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception + +#include +#include +#include +#include +#include +#include + +using namespace std; + +template +void bm(benchmark::State& state) { + vector h(HSize, T{'.'}); + vector n(NSize); + iota(n.begin(), n.end(), T{'a'}); + + static_assert(Pos < HSize); + static_assert(Which < NSize); + h[Pos] = n[Which]; + + for (auto _ : state) { + benchmark::DoNotOptimize(find_first_of(h.begin(), h.end(), n.begin(), n.end())); + } +} + +BENCHMARK(bm); +BENCHMARK(bm); + +BENCHMARK(bm); +BENCHMARK(bm); + +BENCHMARK(bm); +BENCHMARK(bm); + +BENCHMARK(bm); +BENCHMARK(bm); + +BENCHMARK(bm); +BENCHMARK(bm); + +BENCHMARK(bm); +BENCHMARK(bm); + +BENCHMARK_MAIN(); diff --git a/stl/inc/algorithm b/stl/inc/algorithm index 06152ed95db..b00dc4ee874 100644 --- a/stl/inc/algorithm +++ b/stl/inc/algorithm @@ -58,6 +58,11 @@ const void* __stdcall __std_find_last_trivial_2(const void* _First, const void* const void* __stdcall __std_find_last_trivial_4(const void* _First, const void* _Last, uint32_t _Val) noexcept; const void* __stdcall __std_find_last_trivial_8(const void* _First, const void* _Last, uint64_t _Val) noexcept; +const void* __stdcall __std_find_first_of_trivial_1( + const void* _First1, const void* _Last1, const void* _First2, const void* _Last2) noexcept; +const void* __stdcall __std_find_first_of_trivial_2( + const void* _First1, const void* _Last1, const void* _First2, const void* _Last2) noexcept; + __declspec(noalias) _Min_max_1i __stdcall __std_minmax_1i(const void* _First, const void* _Last) noexcept; __declspec(noalias) _Min_max_1u __stdcall __std_minmax_1u(const void* _First, const void* _Last) noexcept; __declspec(noalias) _Min_max_2i __stdcall __std_minmax_2i(const void* _First, const void* _Last) noexcept; @@ -160,6 +165,29 @@ _Ty* __std_find_last_trivial(_Ty* const _First, _Ty* const _Last, const _TVal _V static_assert(_Always_false<_Ty>, "Unexpected size"); } } + +template +_Ty1* __std_find_first_of_trivial( + _Ty1* const _First1, _Ty1* const _Last1, _Ty2* const _First2, _Ty2* const _Last2) noexcept { + if constexpr (sizeof(_Ty1) == 1) { + return const_cast<_Ty1*>( + static_cast(::__std_find_first_of_trivial_1(_First1, _Last1, _First2, _Last2))); + } else if constexpr (sizeof(_Ty1) == 2) { + return const_cast<_Ty1*>( + static_cast(::__std_find_first_of_trivial_2(_First1, _Last1, _First2, _Last2))); + } else { + static_assert(_Always_false<_Ty1>, "Unexpected size"); + } +} + +// find_first_of vectorization is likely to be a win after this size (in elements) +_INLINE_VAR constexpr ptrdiff_t _Threshold_find_first_of = 16; + +// Can we activate the vector algorithms for find_first_of? +template +_INLINE_VAR constexpr bool _Vector_alg_in_find_first_of_is_safe = + _Equal_memcmp_is_safe<_It1, _It2, _Pr> // can replace value comparison with bitwise comparison + && sizeof(_Iter_value_t<_It1>) <= 2; // pcmpestri compatible size _STD_END #endif // _USE_STD_VECTOR_ALGORITHMS @@ -3321,6 +3349,24 @@ _NODISCARD _CONSTEXPR20 _FwdIt1 find_first_of( const auto _ULast1 = _STD _Get_unwrapped(_Last1); const auto _UFirst2 = _STD _Get_unwrapped(_First2); const auto _ULast2 = _STD _Get_unwrapped(_Last2); +#if _USE_STD_VECTOR_ALGORITHMS + if constexpr (_Vector_alg_in_find_first_of_is_safe) { + if (!_STD _Is_constant_evaluated() && _ULast1 - _UFirst1 >= _Threshold_find_first_of) { + const auto _First1_ptr = _STD _To_address(_UFirst1); + const auto _Result = _STD __std_find_first_of_trivial( + _First1_ptr, _STD _To_address(_ULast1), _STD _To_address(_UFirst2), _STD _To_address(_ULast2)); + + if constexpr (is_pointer_v) { + _UFirst1 = _Result; + } else { + _UFirst1 += _Result - _First1_ptr; + } + _STD _Seek_wrapped(_First1, _UFirst1); + return _First1; + } + } +#endif // _USE_STD_VECTOR_ALGORITHMS + for (; _UFirst1 != _ULast1; ++_UFirst1) { for (auto _UMid2 = _UFirst2; _UMid2 != _ULast2; ++_UMid2) { if (_Pred(*_UFirst1, *_UMid2)) { @@ -3398,6 +3444,29 @@ namespace ranges { _STL_INTERNAL_STATIC_ASSERT(sentinel_for<_Se2, _It2>); _STL_INTERNAL_STATIC_ASSERT(indirectly_comparable<_It1, _It2, _Pr, _Pj1, _Pj2>); +#if _USE_STD_VECTOR_ALGORITHMS + if constexpr (_Vector_alg_in_find_first_of_is_safe<_It1, _It2, _Pr> && sized_sentinel_for<_Se1, _It1> + && sized_sentinel_for<_Se2, _It2> && is_same_v<_Pj1, identity> && is_same_v<_Pj2, identity>) { + if (!_STD is_constant_evaluated() && _Last1 - _First1 >= _Threshold_find_first_of) { + const auto _Count1 = _Last1 - _First1; + const auto _First1_ptr = _STD _To_address(_First1); + const auto _Last1_ptr = _First1_ptr + _Count1; + + const auto _Count2 = _Last2 - _First2; + const auto _First2_ptr = _STD _To_address(_First2); + const auto _Last2_ptr = _First2_ptr + _Count2; + + const auto _Result = + _STD __std_find_first_of_trivial(_First1_ptr, _Last1_ptr, _First2_ptr, _Last2_ptr); + + if constexpr (is_pointer_v<_It1>) { + return _Result; + } else { + return _First1 + (_Result - _First1_ptr); + } + } + } +#endif // _USE_STD_VECTOR_ALGORITHMS for (; _First1 != _Last1; ++_First1) { for (auto _Mid2 = _First2; _Mid2 != _Last2; ++_Mid2) { if (_STD invoke(_Pred, _STD invoke(_Proj1, *_First1), _STD invoke(_Proj2, *_Mid2))) { diff --git a/stl/src/vector_algorithms.cpp b/stl/src/vector_algorithms.cpp index d1e2b654e4a..176abaa94e4 100644 --- a/stl/src/vector_algorithms.cpp +++ b/stl/src/vector_algorithms.cpp @@ -2009,6 +2009,74 @@ namespace { } return _Result; } + + template + const void* __stdcall __std_find_first_of_trivial_impl( + const void* _First1, const void* const _Last1, const void* const _First2, const void* const _Last2) noexcept { +#ifndef _M_ARM64EC + const size_t _Needle_length = _Byte_length(_First2, _Last2); + + if (_Use_sse42() && _Needle_length <= 16) { + constexpr int _Op = + (sizeof(_Ty) == 1 ? _SIDD_UBYTE_OPS : _SIDD_UWORD_OPS) | _SIDD_CMP_EQUAL_ANY | _SIDD_LEAST_SIGNIFICANT; + constexpr int _Part_size_el = sizeof(_Ty) == 1 ? 16 : 8; + + const int _Needle_length_el = static_cast(_Needle_length / sizeof(_Ty)); + + alignas(16) uint8_t _Tmp1[16]; + memcpy(_Tmp1, _First2, _Needle_length); + const __m128i _Needle = _mm_load_si128(reinterpret_cast(_Tmp1)); + + const size_t _Haystack_length = _Byte_length(_First1, _Last1); + const void* _Stop_at = _First1; + _Advance_bytes(_Stop_at, _Haystack_length & ~size_t{0xF}); + + while (_First1 != _Stop_at) { + const __m128i _Haystack_part = _mm_loadu_si128(static_cast(_First1)); + + if (_mm_cmpestrc(_Needle, _Needle_length_el, _Haystack_part, _Part_size_el, _Op)) { + const int _Pos = _mm_cmpestri(_Needle, _Needle_length_el, _Haystack_part, _Part_size_el, _Op); + _Advance_bytes(_First1, _Pos * sizeof(_Ty)); + return _First1; + } + + _Advance_bytes(_First1, 16); + } + + const size_t _Last_part_size = _Haystack_length & 0xF; + const int _Last_part_size_el = static_cast(_Last_part_size / sizeof(_Ty)); + + alignas(16) uint8_t _Tmp2[16]; + memcpy(_Tmp2, _First1, _Last_part_size); + const __m128i _Haystack_last_part = _mm_load_si128(reinterpret_cast(_Tmp2)); + + if (_mm_cmpestrc(_Needle, _Needle_length_el, _Haystack_last_part, _Last_part_size_el, _Op)) { + const int _Pos = _mm_cmpestri(_Needle, _Needle_length_el, _Haystack_last_part, _Last_part_size_el, _Op); + _Advance_bytes(_First1, _Pos * sizeof(_Ty)); + return _First1; + } + + _Advance_bytes(_First1, _Last_part_size); + return _First1; + } +#endif // !_M_ARM64EC + + auto _Ptr_haystack = static_cast(_First1); + const auto _Ptr_haystack_end = static_cast(_Last1); + const auto _Ptr_needle = static_cast(_First2); + const auto _Ptr_needle_end = static_cast(_Last2); + + for (; _Ptr_haystack != _Ptr_haystack_end; ++_Ptr_haystack) { + for (auto _Ptr = _Ptr_needle; _Ptr != _Ptr_needle_end; ++_Ptr) { + if (*_Ptr_haystack == *_Ptr) { + return _Ptr_haystack; + } + } + } + + return _Ptr_haystack; + } + } // unnamed namespace extern "C" { @@ -2094,6 +2162,16 @@ __declspec(noalias) size_t return __std_count_trivial_impl<_Find_traits_8>(_First, _Last, _Val); } +const void* __stdcall __std_find_first_of_trivial_1( + const void* _First1, const void* _Last1, const void* _First2, const void* _Last2) noexcept { + return __std_find_first_of_trivial_impl(_First1, _Last1, _First2, _Last2); +} + +const void* __stdcall __std_find_first_of_trivial_2( + const void* _First1, const void* _Last1, const void* _First2, const void* _Last2) noexcept { + return __std_find_first_of_trivial_impl(_First1, _Last1, _First2, _Last2); +} + } // extern "C" #ifndef _M_ARM64EC diff --git a/tests/std/tests/VSO_0000000_vector_algorithms/test.cpp b/tests/std/tests/VSO_0000000_vector_algorithms/test.cpp index 510c7b99353..af4f710cddd 100644 --- a/tests/std/tests/VSO_0000000_vector_algorithms/test.cpp +++ b/tests/std/tests/VSO_0000000_vector_algorithms/test.cpp @@ -120,6 +120,18 @@ auto last_known_good_find_last(FwdIt first, FwdIt last, T v) { } } +template +auto last_known_good_find_first_of(FwdItH h_first, FwdItH h_last, FwdItN n_first, FwdItN n_last) { + for (; h_first != h_last; ++h_first) { + for (FwdItN n = n_first; n != n_last; ++n) { + if (*h_first == *n) { + return h_first; + } + } + } + return h_first; +} + template void test_case_find(const vector& input, T v) { auto expected = last_known_good_find(input.begin(), input.end(), v); @@ -211,6 +223,57 @@ void test_find_last(mt19937_64& gen) { } #endif // _HAS_CXX23 +template +void test_case_find_first_of(const vector& input_haystack, const vector& input_needle) { + auto expected = last_known_good_find_first_of( + input_haystack.begin(), input_haystack.end(), input_needle.begin(), input_needle.end()); + auto actual = find_first_of(input_haystack.begin(), input_haystack.end(), input_needle.begin(), input_needle.end()); + assert(expected == actual); +#if _HAS_CXX20 + auto ranges_actual = ranges::find_first_of(input_haystack, input_needle); + assert(expected == ranges_actual); +#endif // _HAS_CXX20 +} + +template +void test_find_first_of(mt19937_64& gen) { + constexpr size_t needleDataCount = 30; + using TD = conditional_t; + uniform_int_distribution dis('a', 'z'); + vector input_haystack; + vector input_needle; + input_haystack.reserve(dataCount); + input_needle.reserve(needleDataCount); + + for (;;) { + input_needle.clear(); + + test_case_find_first_of(input_haystack, input_needle); + for (size_t attempts = 0; attempts < needleDataCount; ++attempts) { + input_needle.push_back(static_cast(dis(gen))); + test_case_find_first_of(input_haystack, input_needle); + } + + if (input_haystack.size() == dataCount) { + break; + } + + input_haystack.push_back(static_cast(dis(gen))); + } +} + +template +void test_find_first_of_containers() { + C1 haystack{'m', 'e', 'o', 'w', 'C', 'A', 'T', 'S'}; + C2 needle{'R', 'S', 'T'}; + const auto result = find_first_of(haystack.begin(), haystack.end(), needle.begin(), needle.end()); + assert(result == haystack.begin() + 6); +#if _HAS_CXX20 + const auto ranges_result = ranges::find_first_of(haystack, needle); + assert(ranges_result == haystack.begin() + 6); +#endif // _HAS_CXX20 +} + template void test_min_max_element(mt19937_64& gen) { using Limits = numeric_limits; @@ -437,6 +500,24 @@ void test_vector_algorithms(mt19937_64& gen) { test_find_last(gen); #endif // _HAS_CXX23 + test_find_first_of(gen); + test_find_first_of(gen); + test_find_first_of(gen); + test_find_first_of(gen); + test_find_first_of(gen); + test_find_first_of(gen); + test_find_first_of(gen); + test_find_first_of(gen); + test_find_first_of(gen); + + test_find_first_of_containers, vector>(); + test_find_first_of_containers, vector>(); + test_find_first_of_containers, vector>(); + test_find_first_of_containers, const vector>(); + test_find_first_of_containers, const vector>(); + test_find_first_of_containers, vector>(); + test_find_first_of_containers, vector>(); + test_min_max_element(gen); test_min_max_element(gen); test_min_max_element(gen);