From a2c68d2fb98f0b0f909623c3d5ea174d2d8faa92 Mon Sep 17 00:00:00 2001 From: Hari Limaye Date: Wed, 5 Nov 2025 16:54:40 +0000 Subject: [PATCH 1/2] Enable vectorization of std::rotate on ARM64 Add an implementation of _Swap_3_ranges using Neon intrinsics and enable vectorization of std::rotate on ARM64 targets. --- stl/inc/xutility | 2 +- stl/src/vector_algorithms.cpp | 108 +++++++++++++++++++++++++++++++++- 2 files changed, 108 insertions(+), 2 deletions(-) diff --git a/stl/inc/xutility b/stl/inc/xutility index bbaf4357d91..8ba56f690c0 100644 --- a/stl/inc/xutility +++ b/stl/inc/xutility @@ -97,7 +97,7 @@ _STL_DISABLE_CLANG_WARNINGS #define _VECTORIZED_REPLACE _VECTORIZED_FOR_X64_X86 #define _VECTORIZED_REVERSE _VECTORIZED_FOR_X64_X86 #define _VECTORIZED_REVERSE_COPY _VECTORIZED_FOR_X64_X86 -#define _VECTORIZED_ROTATE _VECTORIZED_FOR_X64_X86 +#define _VECTORIZED_ROTATE _VECTORIZED_FOR_X64_X86_ARM64 #define _VECTORIZED_SEARCH _VECTORIZED_FOR_X64_X86 #define _VECTORIZED_SEARCH_N _VECTORIZED_FOR_X64_X86 #define _VECTORIZED_SWAP_RANGES _VECTORIZED_FOR_X64_X86_ARM64 diff --git a/stl/src/vector_algorithms.cpp b/stl/src/vector_algorithms.cpp index ac2d743a32d..42116f41e9b 100644 --- a/stl/src/vector_algorithms.cpp +++ b/stl/src/vector_algorithms.cpp @@ -250,9 +250,113 @@ void* __cdecl __std_swap_ranges_trivially_swappable( } // extern "C" -#ifndef _M_ARM64 namespace { namespace _Rotating { +#ifdef _M_ARM64 + void __forceinline _Swap_3_ranges(void* _First1, void* const _Last1, void* _First2, void* _First3) noexcept { + if (_Byte_length(_First1, _Last1) >= 64) { + constexpr size_t _Mask_64 = ~((static_cast(1) << 6) - 1); + const void* _Stop_at = _First1; + _Advance_bytes(_Stop_at, _Byte_length(_First1, _Last1) & _Mask_64); + do { + const uint8x16_t _Val1Lo1 = vld1q_u8(static_cast(_First1) + 0); + const uint8x16_t _Val1Lo2 = vld1q_u8(static_cast(_First1) + 16); + const uint8x16_t _Val1Hi1 = vld1q_u8(static_cast(_First1) + 32); + const uint8x16_t _Val1Hi2 = vld1q_u8(static_cast(_First1) + 48); + const uint8x16_t _Val2Lo1 = vld1q_u8(static_cast(_First2) + 0); + const uint8x16_t _Val2Lo2 = vld1q_u8(static_cast(_First2) + 16); + const uint8x16_t _Val2Hi1 = vld1q_u8(static_cast(_First2) + 32); + const uint8x16_t _Val2Hi2 = vld1q_u8(static_cast(_First2) + 48); + const uint8x16_t _Val3Lo1 = vld1q_u8(static_cast(_First3) + 0); + const uint8x16_t _Val3Lo2 = vld1q_u8(static_cast(_First3) + 16); + const uint8x16_t _Val3Hi1 = vld1q_u8(static_cast(_First3) + 32); + const uint8x16_t _Val3Hi2 = vld1q_u8(static_cast(_First3) + 48); + vst1q_u8(static_cast(_First1) + 0, _Val2Lo1); + vst1q_u8(static_cast(_First1) + 16, _Val2Lo2); + vst1q_u8(static_cast(_First1) + 32, _Val2Hi1); + vst1q_u8(static_cast(_First1) + 48, _Val2Hi2); + vst1q_u8(static_cast(_First2) + 0, _Val3Lo1); + vst1q_u8(static_cast(_First2) + 16, _Val3Lo2); + vst1q_u8(static_cast(_First2) + 32, _Val3Hi1); + vst1q_u8(static_cast(_First2) + 48, _Val3Hi2); + vst1q_u8(static_cast(_First3) + 0, _Val1Lo1); + vst1q_u8(static_cast(_First3) + 16, _Val1Lo2); + vst1q_u8(static_cast(_First3) + 32, _Val1Hi1); + vst1q_u8(static_cast(_First3) + 48, _Val1Hi2); + _Advance_bytes(_First1, 64); + _Advance_bytes(_First2, 64); + _Advance_bytes(_First3, 64); + } while (_First1 != _Stop_at); + } + + if (_Byte_length(_First1, _Last1) >= 32) { + const uint8x16_t _Val1Lo = vld1q_u8(static_cast(_First1) + 0); + const uint8x16_t _Val1Hi = vld1q_u8(static_cast(_First1) + 16); + const uint8x16_t _Val2Lo = vld1q_u8(static_cast(_First2) + 0); + const uint8x16_t _Val2Hi = vld1q_u8(static_cast(_First2) + 16); + const uint8x16_t _Val3Lo = vld1q_u8(static_cast(_First3) + 0); + const uint8x16_t _Val3Hi = vld1q_u8(static_cast(_First3) + 16); + vst1q_u8(static_cast(_First1) + 0, _Val2Lo); + vst1q_u8(static_cast(_First1) + 16, _Val2Hi); + vst1q_u8(static_cast(_First2) + 0, _Val3Lo); + vst1q_u8(static_cast(_First2) + 16, _Val3Hi); + vst1q_u8(static_cast(_First3) + 0, _Val1Lo); + vst1q_u8(static_cast(_First3) + 16, _Val1Hi); + _Advance_bytes(_First1, 32); + _Advance_bytes(_First2, 32); + _Advance_bytes(_First3, 32); + } + + if (_Byte_length(_First1, _Last1) >= 16) { + const uint8x16_t _Val1 = vld1q_u8(static_cast(_First1)); + const uint8x16_t _Val2 = vld1q_u8(static_cast(_First2)); + const uint8x16_t _Val3 = vld1q_u8(static_cast(_First3)); + vst1q_u8(static_cast(_First1), _Val2); + vst1q_u8(static_cast(_First2), _Val3); + vst1q_u8(static_cast(_First3), _Val1); + _Advance_bytes(_First1, 16); + _Advance_bytes(_First2, 16); + _Advance_bytes(_First3, 16); + } + + if (_Byte_length(_First1, _Last1) >= 8) { + const uint8x8_t _Val1 = vld1_u8(static_cast(_First1)); + const uint8x8_t _Val2 = vld1_u8(static_cast(_First2)); + const uint8x8_t _Val3 = vld1_u8(static_cast(_First3)); + vst1_u8(static_cast(_First1), _Val2); + vst1_u8(static_cast(_First2), _Val3); + vst1_u8(static_cast(_First3), _Val1); + _Advance_bytes(_First1, 8); + _Advance_bytes(_First2, 8); + _Advance_bytes(_First3, 8); + } + + if (_Byte_length(_First1, _Last1) >= 4) { + uint32x2_t _Val1 = vdup_n_u32(0); + uint32x2_t _Val2 = vdup_n_u32(0); + uint32x2_t _Val3 = vdup_n_u32(0); + _Val1 = vld1_lane_u32(static_cast(_First1), _Val1, 0); + _Val2 = vld1_lane_u32(static_cast(_First2), _Val2, 0); + _Val3 = vld1_lane_u32(static_cast(_First3), _Val3, 0); + vst1_lane_u32(static_cast(_First1), _Val2, 0); + vst1_lane_u32(static_cast(_First2), _Val3, 0); + vst1_lane_u32(static_cast(_First3), _Val1, 0); + _Advance_bytes(_First1, 4); + _Advance_bytes(_First2, 4); + _Advance_bytes(_First3, 4); + } + + auto _First1c = static_cast(_First1); + auto _First2c = static_cast(_First2); + auto _First3c = static_cast(_First3); + for (; _First1c != _Last1; ++_First1c, ++_First2c, ++_First3c) { + const unsigned char _Ch = *_First1c; + *_First1c = *_First2c; + *_First2c = *_First3c; + *_First3c = _Ch; + } + } +#else // ^^^ defined(_M_ARM64) / !defined(_M_ARM64) vvv void _Swap_3_ranges(void* _First1, void* const _Last1, void* _First2, void* _First3) noexcept { #ifndef _M_ARM64EC constexpr size_t _Mask_32 = ~((static_cast(1) << 5) - 1); @@ -346,6 +450,7 @@ namespace { *_First3c = _Ch; } } +#endif // ^^^ !defined(_M_ARM64) ^^^ constexpr size_t _Buf_size = 512; @@ -418,6 +523,7 @@ __declspec(noalias) void __stdcall __std_rotate(void* _First, void* const _Mid, } // extern "C" +#ifndef _M_ARM64 namespace { namespace _Reversing { #ifdef _M_ARM64EC From 77d8aa17823f5f480db5ea4a1aae380019390811 Mon Sep 17 00:00:00 2001 From: Hari Limaye Date: Tue, 11 Nov 2025 13:01:30 +0000 Subject: [PATCH 2/2] Fix formatting --- stl/src/vector_algorithms.cpp | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/stl/src/vector_algorithms.cpp b/stl/src/vector_algorithms.cpp index 42116f41e9b..9b78d5cbd72 100644 --- a/stl/src/vector_algorithms.cpp +++ b/stl/src/vector_algorithms.cpp @@ -256,7 +256,7 @@ namespace { void __forceinline _Swap_3_ranges(void* _First1, void* const _Last1, void* _First2, void* _First3) noexcept { if (_Byte_length(_First1, _Last1) >= 64) { constexpr size_t _Mask_64 = ~((static_cast(1) << 6) - 1); - const void* _Stop_at = _First1; + const void* _Stop_at = _First1; _Advance_bytes(_Stop_at, _Byte_length(_First1, _Last1) & _Mask_64); do { const uint8x16_t _Val1Lo1 = vld1q_u8(static_cast(_First1) + 0); @@ -332,12 +332,12 @@ namespace { } if (_Byte_length(_First1, _Last1) >= 4) { - uint32x2_t _Val1 = vdup_n_u32(0); - uint32x2_t _Val2 = vdup_n_u32(0); - uint32x2_t _Val3 = vdup_n_u32(0); - _Val1 = vld1_lane_u32(static_cast(_First1), _Val1, 0); - _Val2 = vld1_lane_u32(static_cast(_First2), _Val2, 0); - _Val3 = vld1_lane_u32(static_cast(_First3), _Val3, 0); + uint32x2_t _Val1 = vdup_n_u32(0); + uint32x2_t _Val2 = vdup_n_u32(0); + uint32x2_t _Val3 = vdup_n_u32(0); + _Val1 = vld1_lane_u32(static_cast(_First1), _Val1, 0); + _Val2 = vld1_lane_u32(static_cast(_First2), _Val2, 0); + _Val3 = vld1_lane_u32(static_cast(_First3), _Val3, 0); vst1_lane_u32(static_cast(_First1), _Val2, 0); vst1_lane_u32(static_cast(_First2), _Val3, 0); vst1_lane_u32(static_cast(_First3), _Val1, 0);