diff options
| author | Raghuveer Devulapalli <raghuveer.devulapalli@intel.com> | 2022-09-19 10:35:03 -0700 |
|---|---|---|
| committer | Raghuveer Devulapalli <raghuveer.devulapalli@intel.com> | 2023-01-30 13:38:39 -0800 |
| commit | 49278b961b7254bc6a4aee478587c69682a3827e (patch) | |
| tree | 9bcf5e3f022df96018d2dcd158578c16218b829f /numpy/core/src | |
| parent | c662a712a30b1b640a80421619bb97556ffe965b (diff) | |
| download | numpy-49278b961b7254bc6a4aee478587c69682a3827e.tar.gz | |
ENH: Add x86-simd-sort source files
Diffstat (limited to 'numpy/core/src')
4 files changed, 2277 insertions, 0 deletions
diff --git a/numpy/core/src/npysort/x86-simd-sort/src/avx512-16bit-qsort.hpp b/numpy/core/src/npysort/x86-simd-sort/src/avx512-16bit-qsort.hpp new file mode 100644 index 000000000..1673eb5da --- /dev/null +++ b/numpy/core/src/npysort/x86-simd-sort/src/avx512-16bit-qsort.hpp @@ -0,0 +1,527 @@ +/******************************************************************* + * Copyright (C) 2022 Intel Corporation + * SPDX-License-Identifier: BSD-3-Clause + * Authors: Raghuveer Devulapalli <raghuveer.devulapalli@intel.com> + * ****************************************************************/ + +#ifndef __AVX512_QSORT_16BIT__ +#define __AVX512_QSORT_16BIT__ + +#include "avx512-common-qsort.h" + +/* + * Constants used in sorting 32 elements in a ZMM registers. Based on Bitonic + * sorting network (see + * https://en.wikipedia.org/wiki/Bitonic_sorter#/media/File:BitonicSort.svg) + */ +// ZMM register: 31,30,29,28,27,26,25,24,23,22,21,20,19,18,17,16,15,14,13,12,11,10,9,8,7,6,5,4,3,2,1,0 +#define NETWORK_16BIT_1 \ + 24, 25, 26, 27, 28, 29, 30, 31, 16, 17, 18, 19, 20, 21, 22, 23, 8, 9, 10, \ + 11, 12, 13, 14, 15, 0, 1, 2, 3, 4, 5, 6, 7 +#define NETWORK_16BIT_2 \ + 16, 17, 18, 19, 20, 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31, 0, 1, 2, \ + 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15 +#define NETWORK_16BIT_3 \ + 27, 26, 25, 24, 31, 30, 29, 28, 19, 18, 17, 16, 23, 22, 21, 20, 11, 10, 9, \ + 8, 15, 14, 13, 12, 3, 2, 1, 0, 7, 6, 5, 4 +#define NETWORK_16BIT_4 \ + 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15, 16, 17, 18, 19, 20, \ + 21, 22, 23, 24, 25, 26, 27, 28, 29, 30, 31 +#define NETWORK_16BIT_5 \ + 23, 22, 21, 20, 19, 18, 17, 16, 31, 30, 29, 28, 27, 26, 25, 24, 7, 6, 5, \ + 4, 3, 2, 1, 0, 15, 14, 13, 12, 11, 10, 9, 8 +#define NETWORK_16BIT_6 \ + 15, 14, 13, 12, 11, 10, 9, 8, 7, 6, 5, 4, 3, 2, 1, 0, 31, 30, 29, 28, 27, \ + 26, 25, 24, 23, 22, 21, 20, 19, 18, 17, 16 + +template <> +struct vector<int16_t> { + using type_t = int16_t; + using zmm_t = __m512i; + using ymm_t = __m256i; + using opmask_t = __mmask32; + static const uint8_t numlanes = 32; + + static type_t type_max() + { + return X86_SIMD_SORT_MAX_INT16; + } + static type_t type_min() + { + return X86_SIMD_SORT_MIN_INT16; + } + static zmm_t zmm_max() + { + return _mm512_set1_epi16(type_max()); + } + + static opmask_t knot_opmask(opmask_t x) + { + return _knot_mask32(x); + } + static opmask_t ge(zmm_t x, zmm_t y) + { + return _mm512_cmp_epi16_mask(x, y, _MM_CMPINT_NLT); + } + //template <int scale> + //static zmm_t i64gather(__m512i index, void const *base) + //{ + // return _mm512_i64gather_epi64(index, base, scale); + //} + static zmm_t loadu(void const *mem) + { + return _mm512_loadu_si512(mem); + } + static zmm_t max(zmm_t x, zmm_t y) + { + return _mm512_max_epi16(x, y); + } + static void mask_compressstoreu(void *mem, opmask_t mask, zmm_t x) + { + // AVX512_VBMI2 + return _mm512_mask_compressstoreu_epi16(mem, mask, x); + } + static zmm_t mask_loadu(zmm_t x, opmask_t mask, void const *mem) + { + // AVX512BW + return _mm512_mask_loadu_epi16(x, mask, mem); + } + static zmm_t mask_mov(zmm_t x, opmask_t mask, zmm_t y) + { + return _mm512_mask_mov_epi16(x, mask, y); + } + static void mask_storeu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_storeu_epi16(mem, mask, x); + } + static zmm_t min(zmm_t x, zmm_t y) + { + return _mm512_min_epi16(x, y); + } + static zmm_t permutexvar(__m512i idx, zmm_t zmm) + { + return _mm512_permutexvar_epi16(idx, zmm); + } + static type_t reducemax(zmm_t v) + { + zmm_t lo = _mm512_cvtepi16_epi32(_mm512_extracti64x4_epi64(v, 0)); + zmm_t hi = _mm512_cvtepi16_epi32(_mm512_extracti64x4_epi64(v, 1)); + type_t lo_max = (type_t)_mm512_reduce_max_epi32(lo); + type_t hi_max = (type_t)_mm512_reduce_max_epi32(hi); + return std::max(lo_max, hi_max); + } + static type_t reducemin(zmm_t v) + { + zmm_t lo = _mm512_cvtepi16_epi32(_mm512_extracti64x4_epi64(v, 0)); + zmm_t hi = _mm512_cvtepi16_epi32(_mm512_extracti64x4_epi64(v, 1)); + type_t lo_min = (type_t)_mm512_reduce_min_epi32(lo); + type_t hi_min = (type_t)_mm512_reduce_min_epi32(hi); + return std::min(lo_min, hi_min); + } + static zmm_t set1(type_t v) + { + return _mm512_set1_epi16(v); + } + template <uint8_t mask> + static zmm_t shuffle(zmm_t zmm) + { + zmm = _mm512_shufflehi_epi16(zmm, (_MM_PERM_ENUM)mask); + return _mm512_shufflelo_epi16(zmm, (_MM_PERM_ENUM)mask); + } + static void storeu(void *mem, zmm_t x) + { + return _mm512_storeu_si512(mem, x); + } +}; +template <> +struct vector<uint16_t> { + using type_t = uint16_t; + using zmm_t = __m512i; + using ymm_t = __m256i; + using opmask_t = __mmask32; + static const uint8_t numlanes = 32; + + static type_t type_max() + { + return X86_SIMD_SORT_MAX_UINT16; + } + static type_t type_min() + { + return 0; + } + static zmm_t zmm_max() + { + return _mm512_set1_epi16(type_max()); + } // TODO: this should broadcast bits as is? + + //template<int scale> + //static zmm_t i64gather(__m512i index, void const *base) + //{ + // return _mm512_i64gather_epi64(index, base, scale); + //} + static opmask_t knot_opmask(opmask_t x) + { + return _knot_mask32(x); + } + static opmask_t ge(zmm_t x, zmm_t y) + { + return _mm512_cmp_epu16_mask(x, y, _MM_CMPINT_NLT); + } + static zmm_t loadu(void const *mem) + { + return _mm512_loadu_si512(mem); + } + static zmm_t max(zmm_t x, zmm_t y) + { + return _mm512_max_epu16(x, y); + } + static void mask_compressstoreu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_compressstoreu_epi16(mem, mask, x); + } + static zmm_t mask_loadu(zmm_t x, opmask_t mask, void const *mem) + { + return _mm512_mask_loadu_epi16(x, mask, mem); + } + static zmm_t mask_mov(zmm_t x, opmask_t mask, zmm_t y) + { + return _mm512_mask_mov_epi16(x, mask, y); + } + static void mask_storeu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_storeu_epi16(mem, mask, x); + } + static zmm_t min(zmm_t x, zmm_t y) + { + return _mm512_min_epu16(x, y); + } + static zmm_t permutexvar(__m512i idx, zmm_t zmm) + { + return _mm512_permutexvar_epi16(idx, zmm); + } + static type_t reducemax(zmm_t v) + { + zmm_t lo = _mm512_cvtepu16_epi32(_mm512_extracti64x4_epi64(v, 0)); + zmm_t hi = _mm512_cvtepu16_epi32(_mm512_extracti64x4_epi64(v, 1)); + type_t lo_max = (type_t)_mm512_reduce_max_epi32(lo); + type_t hi_max = (type_t)_mm512_reduce_max_epi32(hi); + return std::max(lo_max, hi_max); + } + static type_t reducemin(zmm_t v) + { + zmm_t lo = _mm512_cvtepu16_epi32(_mm512_extracti64x4_epi64(v, 0)); + zmm_t hi = _mm512_cvtepu16_epi32(_mm512_extracti64x4_epi64(v, 1)); + type_t lo_min = (type_t)_mm512_reduce_min_epi32(lo); + type_t hi_min = (type_t)_mm512_reduce_min_epi32(hi); + return std::min(lo_min, hi_min); + } + static zmm_t set1(type_t v) + { + return _mm512_set1_epi16(v); + } + template <uint8_t mask> + static zmm_t shuffle(zmm_t zmm) + { + zmm = _mm512_shufflehi_epi16(zmm, (_MM_PERM_ENUM)mask); + return _mm512_shufflelo_epi16(zmm, (_MM_PERM_ENUM)mask); + } + static void storeu(void *mem, zmm_t x) + { + return _mm512_storeu_si512(mem, x); + } +}; + +/* + * Assumes zmm is random and performs a full sorting network defined in + * https://en.wikipedia.org/wiki/Bitonic_sorter#/media/File:BitonicSort.svg + */ +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline zmm_t sort_zmm_16bit(zmm_t zmm) +{ + // Level 1 + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAAAAAA); + // Level 2 + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(0, 1, 2, 3)>(zmm), + 0xCCCCCCCC); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAAAAAA); + // Level 3 + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_1), zmm), + 0xF0F0F0F0); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(1, 0, 3, 2)>(zmm), + 0xCCCCCCCC); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAAAAAA); + // Level 4 + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_2), zmm), + 0xFF00FF00); + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_3), zmm), + 0xF0F0F0F0); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(1, 0, 3, 2)>(zmm), + 0xCCCCCCCC); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAAAAAA); + // Level 5 + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_4), zmm), + 0xFFFF0000); + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_5), zmm), + 0xFF00FF00); + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_3), zmm), + 0xF0F0F0F0); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(1, 0, 3, 2)>(zmm), + 0xCCCCCCCC); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAAAAAA); + return zmm; +} + +// Assumes zmm is bitonic and performs a recursive half cleaner +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline zmm_t bitonic_merge_zmm_16bit(zmm_t zmm) +{ + // 1) half_cleaner[32]: compare 1-17, 2-18, 3-19 etc .. + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_6), zmm), + 0xFFFF0000); + // 2) half_cleaner[16]: compare 1-9, 2-10, 3-11 etc .. + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_5), zmm), + 0xFF00FF00); + // 3) half_cleaner[8] + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_3), zmm), + 0xF0F0F0F0); + // 3) half_cleaner[4] + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(1, 0, 3, 2)>(zmm), + 0xCCCCCCCC); + // 3) half_cleaner[2] + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAAAAAA); + return zmm; +} + +// Assumes zmm1 and zmm2 are sorted and performs a recursive half cleaner +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline void bitonic_merge_two_zmm_16bit(zmm_t &zmm1, zmm_t &zmm2) +{ + // 1) First step of a merging network: coex of zmm1 and zmm2 reversed + zmm2 = vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_4), zmm2); + zmm_t zmm3 = vtype::min(zmm1, zmm2); + zmm_t zmm4 = vtype::max(zmm1, zmm2); + // 2) Recursive half cleaner for each + zmm1 = bitonic_merge_zmm_16bit<vtype>(zmm3); + zmm2 = bitonic_merge_zmm_16bit<vtype>(zmm4); +} + +// Assumes [zmm0, zmm1] and [zmm2, zmm3] are sorted and performs a recursive +// half cleaner +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline void bitonic_merge_four_zmm_16bit(zmm_t *zmm) +{ + zmm_t zmm2r = vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_4), zmm[2]); + zmm_t zmm3r = vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_4), zmm[3]); + zmm_t zmm_t1 = vtype::min(zmm[0], zmm3r); + zmm_t zmm_t2 = vtype::min(zmm[1], zmm2r); + zmm_t zmm_t3 = vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_4), + vtype::max(zmm[1], zmm2r)); + zmm_t zmm_t4 = vtype::permutexvar(_mm512_set_epi16(NETWORK_16BIT_4), + vtype::max(zmm[0], zmm3r)); + zmm_t zmm0 = vtype::min(zmm_t1, zmm_t2); + zmm_t zmm1 = vtype::max(zmm_t1, zmm_t2); + zmm_t zmm2 = vtype::min(zmm_t3, zmm_t4); + zmm_t zmm3 = vtype::max(zmm_t3, zmm_t4); + zmm[0] = bitonic_merge_zmm_16bit<vtype>(zmm0); + zmm[1] = bitonic_merge_zmm_16bit<vtype>(zmm1); + zmm[2] = bitonic_merge_zmm_16bit<vtype>(zmm2); + zmm[3] = bitonic_merge_zmm_16bit<vtype>(zmm3); +} + +template <typename vtype, typename type_t> +static inline void sort_32_16bit(type_t *arr, int32_t N) +{ + typename vtype::opmask_t load_mask = ((0x1ull << N) - 0x1ull) & 0xFFFFFFFF; + typename vtype::zmm_t zmm + = vtype::mask_loadu(vtype::zmm_max(), load_mask, arr); + vtype::mask_storeu(arr, load_mask, sort_zmm_16bit<vtype>(zmm)); +} + +template <typename vtype, typename type_t> +static inline void sort_64_16bit(type_t *arr, int32_t N) +{ + if (N <= 32) { + sort_32_16bit<vtype>(arr, N); + return; + } + using zmm_t = typename vtype::zmm_t; + typename vtype::opmask_t load_mask + = ((0x1ull << (N - 32)) - 0x1ull) & 0xFFFFFFFF; + zmm_t zmm1 = vtype::loadu(arr); + zmm_t zmm2 = vtype::mask_loadu(vtype::zmm_max(), load_mask, arr + 32); + zmm1 = sort_zmm_16bit<vtype>(zmm1); + zmm2 = sort_zmm_16bit<vtype>(zmm2); + bitonic_merge_two_zmm_16bit<vtype>(zmm1, zmm2); + vtype::storeu(arr, zmm1); + vtype::mask_storeu(arr + 32, load_mask, zmm2); +} + +template <typename vtype, typename type_t> +static inline void sort_128_16bit(type_t *arr, int32_t N) +{ + if (N <= 64) { + sort_64_16bit<vtype>(arr, N); + return; + } + using zmm_t = typename vtype::zmm_t; + using opmask_t = typename vtype::opmask_t; + zmm_t zmm[4]; + zmm[0] = vtype::loadu(arr); + zmm[1] = vtype::loadu(arr + 32); + opmask_t load_mask1 = 0xFFFFFFFF, load_mask2 = 0xFFFFFFFF; + if (N != 128) { + uint64_t combined_mask = (0x1ull << (N - 64)) - 0x1ull; + load_mask1 = combined_mask & 0xFFFFFFFF; + load_mask2 = (combined_mask >> 32) & 0xFFFFFFFF; + } + zmm[2] = vtype::mask_loadu(vtype::zmm_max(), load_mask1, arr + 64); + zmm[3] = vtype::mask_loadu(vtype::zmm_max(), load_mask2, arr + 96); + zmm[0] = sort_zmm_16bit<vtype>(zmm[0]); + zmm[1] = sort_zmm_16bit<vtype>(zmm[1]); + zmm[2] = sort_zmm_16bit<vtype>(zmm[2]); + zmm[3] = sort_zmm_16bit<vtype>(zmm[3]); + bitonic_merge_two_zmm_16bit<vtype>(zmm[0], zmm[1]); + bitonic_merge_two_zmm_16bit<vtype>(zmm[2], zmm[3]); + bitonic_merge_four_zmm_16bit<vtype>(zmm); + vtype::storeu(arr, zmm[0]); + vtype::storeu(arr + 32, zmm[1]); + vtype::mask_storeu(arr + 64, load_mask1, zmm[2]); + vtype::mask_storeu(arr + 96, load_mask2, zmm[3]); +} + +template <typename vtype, typename type_t> +static inline type_t +get_pivot_16bit(type_t *arr, const int64_t left, const int64_t right) +{ + // median of 32 + int64_t size = (right - left) / 32; + __m512i rand_vec = _mm512_set_epi16(arr[left], + arr[left + size], + arr[left + 2 * size], + arr[left + 3 * size], + arr[left + 4 * size], + arr[left + 5 * size], + arr[left + 6 * size], + arr[left + 7 * size], + arr[left + 8 * size], + arr[left + 9 * size], + arr[left + 10 * size], + arr[left + 11 * size], + arr[left + 12 * size], + arr[left + 13 * size], + arr[left + 14 * size], + arr[left + 15 * size], + arr[left + 16 * size], + arr[left + 17 * size], + arr[left + 18 * size], + arr[left + 19 * size], + arr[left + 20 * size], + arr[left + 21 * size], + arr[left + 22 * size], + arr[left + 23 * size], + arr[left + 24 * size], + arr[left + 25 * size], + arr[left + 26 * size], + arr[left + 27 * size], + arr[left + 28 * size], + arr[left + 29 * size], + arr[left + 30 * size], + arr[left + 31 * size]); + __m512i sort = sort_zmm_16bit<vtype>(rand_vec); + return ((type_t *)&sort)[16]; +} + +template <typename vtype, typename type_t> +static inline void +qsort_16bit_(type_t *arr, int64_t left, int64_t right, int64_t max_iters) +{ + /* + * Resort to std::sort if quicksort isnt making any progress + */ + if (max_iters <= 0) { + std::sort(arr + left, arr + right + 1); + return; + } + /* + * Base case: use bitonic networks to sort arrays <= 128 + */ + if (right + 1 - left <= 128) { + sort_128_16bit<vtype>(arr + left, (int32_t)(right + 1 - left)); + return; + } + + type_t pivot = get_pivot_16bit<vtype>(arr, left, right); + type_t smallest = vtype::type_max(); + type_t biggest = vtype::type_min(); + int64_t pivot_index = partition_avx512<vtype>( + arr, left, right + 1, pivot, &smallest, &biggest); + if (pivot != smallest) + qsort_16bit_<vtype>(arr, left, pivot_index - 1, max_iters - 1); + if (pivot != biggest) + qsort_16bit_<vtype>(arr, pivot_index, right, max_iters - 1); +} + +template <> +void avx512_qsort(int16_t *arr, int64_t arrsize) +{ + if (arrsize > 1) { + qsort_16bit_<vector<int16_t>, int16_t>( + arr, 0, arrsize - 1, 2 * (int64_t)log2(arrsize)); + } +} + +template <> +void avx512_qsort(uint16_t *arr, int64_t arrsize) +{ + if (arrsize > 1) { + qsort_16bit_<vector<uint16_t>, uint16_t>( + arr, 0, arrsize - 1, 2 * (int64_t)log2(arrsize)); + } +} +#endif // __AVX512_QSORT_16BIT__ diff --git a/numpy/core/src/npysort/x86-simd-sort/src/avx512-32bit-qsort.hpp b/numpy/core/src/npysort/x86-simd-sort/src/avx512-32bit-qsort.hpp new file mode 100644 index 000000000..cbc5368f0 --- /dev/null +++ b/numpy/core/src/npysort/x86-simd-sort/src/avx512-32bit-qsort.hpp @@ -0,0 +1,712 @@ +/******************************************************************* + * Copyright (C) 2022 Intel Corporation + * Copyright (C) 2021 Serge Sans Paille + * SPDX-License-Identifier: BSD-3-Clause + * Authors: Raghuveer Devulapalli <raghuveer.devulapalli@intel.com> + * Serge Sans Paille <serge.guelton@telecom-bretagne.eu> + * ****************************************************************/ +#ifndef __AVX512_QSORT_32BIT__ +#define __AVX512_QSORT_32BIT__ + +#include "avx512-common-qsort.h" + +/* + * Constants used in sorting 16 elements in a ZMM registers. Based on Bitonic + * sorting network (see + * https://en.wikipedia.org/wiki/Bitonic_sorter#/media/File:BitonicSort.svg) + */ +#define NETWORK_32BIT_1 14, 15, 12, 13, 10, 11, 8, 9, 6, 7, 4, 5, 2, 3, 0, 1 +#define NETWORK_32BIT_2 12, 13, 14, 15, 8, 9, 10, 11, 4, 5, 6, 7, 0, 1, 2, 3 +#define NETWORK_32BIT_3 8, 9, 10, 11, 12, 13, 14, 15, 0, 1, 2, 3, 4, 5, 6, 7 +#define NETWORK_32BIT_4 13, 12, 15, 14, 9, 8, 11, 10, 5, 4, 7, 6, 1, 0, 3, 2 +#define NETWORK_32BIT_5 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12, 13, 14, 15 +#define NETWORK_32BIT_6 11, 10, 9, 8, 15, 14, 13, 12, 3, 2, 1, 0, 7, 6, 5, 4 +#define NETWORK_32BIT_7 7, 6, 5, 4, 3, 2, 1, 0, 15, 14, 13, 12, 11, 10, 9, 8 + +template <> +struct vector<int32_t> { + using type_t = int32_t; + using zmm_t = __m512i; + using ymm_t = __m256i; + using opmask_t = __mmask16; + static const uint8_t numlanes = 16; + + static type_t type_max() + { + return X86_SIMD_SORT_MAX_INT32; + } + static type_t type_min() + { + return X86_SIMD_SORT_MIN_INT32; + } + static zmm_t zmm_max() + { + return _mm512_set1_epi32(type_max()); + } + + static opmask_t knot_opmask(opmask_t x) + { + return _knot_mask16(x); + } + static opmask_t ge(zmm_t x, zmm_t y) + { + return _mm512_cmp_epi32_mask(x, y, _MM_CMPINT_NLT); + } + template <int scale> + static ymm_t i64gather(__m512i index, void const *base) + { + return _mm512_i64gather_epi32(index, base, scale); + } + static zmm_t merge(ymm_t y1, ymm_t y2) + { + zmm_t z1 = _mm512_castsi256_si512(y1); + return _mm512_inserti32x8(z1, y2, 1); + } + static zmm_t loadu(void const *mem) + { + return _mm512_loadu_si512(mem); + } + static void mask_compressstoreu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_compressstoreu_epi32(mem, mask, x); + } + static zmm_t mask_loadu(zmm_t x, opmask_t mask, void const *mem) + { + return _mm512_mask_loadu_epi32(x, mask, mem); + } + static zmm_t mask_mov(zmm_t x, opmask_t mask, zmm_t y) + { + return _mm512_mask_mov_epi32(x, mask, y); + } + static void mask_storeu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_storeu_epi32(mem, mask, x); + } + static zmm_t min(zmm_t x, zmm_t y) + { + return _mm512_min_epi32(x, y); + } + static zmm_t max(zmm_t x, zmm_t y) + { + return _mm512_max_epi32(x, y); + } + static zmm_t permutexvar(__m512i idx, zmm_t zmm) + { + return _mm512_permutexvar_epi32(idx, zmm); + } + static type_t reducemax(zmm_t v) + { + return _mm512_reduce_max_epi32(v); + } + static type_t reducemin(zmm_t v) + { + return _mm512_reduce_min_epi32(v); + } + static zmm_t set1(type_t v) + { + return _mm512_set1_epi32(v); + } + template <uint8_t mask> + static zmm_t shuffle(zmm_t zmm) + { + return _mm512_shuffle_epi32(zmm, (_MM_PERM_ENUM)mask); + } + static void storeu(void *mem, zmm_t x) + { + return _mm512_storeu_si512(mem, x); + } + + static ymm_t max(ymm_t x, ymm_t y) + { + return _mm256_max_epi32(x, y); + } + static ymm_t min(ymm_t x, ymm_t y) + { + return _mm256_min_epi32(x, y); + } +}; +template <> +struct vector<uint32_t> { + using type_t = uint32_t; + using zmm_t = __m512i; + using ymm_t = __m256i; + using opmask_t = __mmask16; + static const uint8_t numlanes = 16; + + static type_t type_max() + { + return X86_SIMD_SORT_MAX_UINT32; + } + static type_t type_min() + { + return 0; + } + static zmm_t zmm_max() + { + return _mm512_set1_epi32(type_max()); + } // TODO: this should broadcast bits as is? + + template <int scale> + static ymm_t i64gather(__m512i index, void const *base) + { + return _mm512_i64gather_epi32(index, base, scale); + } + static zmm_t merge(ymm_t y1, ymm_t y2) + { + zmm_t z1 = _mm512_castsi256_si512(y1); + return _mm512_inserti32x8(z1, y2, 1); + } + static opmask_t knot_opmask(opmask_t x) + { + return _knot_mask16(x); + } + static opmask_t ge(zmm_t x, zmm_t y) + { + return _mm512_cmp_epu32_mask(x, y, _MM_CMPINT_NLT); + } + static zmm_t loadu(void const *mem) + { + return _mm512_loadu_si512(mem); + } + static zmm_t max(zmm_t x, zmm_t y) + { + return _mm512_max_epu32(x, y); + } + static void mask_compressstoreu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_compressstoreu_epi32(mem, mask, x); + } + static zmm_t mask_loadu(zmm_t x, opmask_t mask, void const *mem) + { + return _mm512_mask_loadu_epi32(x, mask, mem); + } + static zmm_t mask_mov(zmm_t x, opmask_t mask, zmm_t y) + { + return _mm512_mask_mov_epi32(x, mask, y); + } + static void mask_storeu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_storeu_epi32(mem, mask, x); + } + static zmm_t min(zmm_t x, zmm_t y) + { + return _mm512_min_epu32(x, y); + } + static zmm_t permutexvar(__m512i idx, zmm_t zmm) + { + return _mm512_permutexvar_epi32(idx, zmm); + } + static type_t reducemax(zmm_t v) + { + return _mm512_reduce_max_epu32(v); + } + static type_t reducemin(zmm_t v) + { + return _mm512_reduce_min_epu32(v); + } + static zmm_t set1(type_t v) + { + return _mm512_set1_epi32(v); + } + template <uint8_t mask> + static zmm_t shuffle(zmm_t zmm) + { + return _mm512_shuffle_epi32(zmm, (_MM_PERM_ENUM)mask); + } + static void storeu(void *mem, zmm_t x) + { + return _mm512_storeu_si512(mem, x); + } + + static ymm_t max(ymm_t x, ymm_t y) + { + return _mm256_max_epu32(x, y); + } + static ymm_t min(ymm_t x, ymm_t y) + { + return _mm256_min_epu32(x, y); + } +}; +template <> +struct vector<float> { + using type_t = float; + using zmm_t = __m512; + using ymm_t = __m256; + using opmask_t = __mmask16; + static const uint8_t numlanes = 16; + + static type_t type_max() + { + return X86_SIMD_SORT_INFINITYF; + } + static type_t type_min() + { + return -X86_SIMD_SORT_INFINITYF; + } + static zmm_t zmm_max() + { + return _mm512_set1_ps(type_max()); + } + + static opmask_t knot_opmask(opmask_t x) + { + return _knot_mask16(x); + } + static opmask_t ge(zmm_t x, zmm_t y) + { + return _mm512_cmp_ps_mask(x, y, _CMP_GE_OQ); + } + template <int scale> + static ymm_t i64gather(__m512i index, void const *base) + { + return _mm512_i64gather_ps(index, base, scale); + } + static zmm_t merge(ymm_t y1, ymm_t y2) + { + zmm_t z1 = _mm512_castsi512_ps( + _mm512_castsi256_si512(_mm256_castps_si256(y1))); + return _mm512_insertf32x8(z1, y2, 1); + } + static zmm_t loadu(void const *mem) + { + return _mm512_loadu_ps(mem); + } + static zmm_t max(zmm_t x, zmm_t y) + { + return _mm512_max_ps(x, y); + } + static void mask_compressstoreu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_compressstoreu_ps(mem, mask, x); + } + static zmm_t mask_loadu(zmm_t x, opmask_t mask, void const *mem) + { + return _mm512_mask_loadu_ps(x, mask, mem); + } + static zmm_t mask_mov(zmm_t x, opmask_t mask, zmm_t y) + { + return _mm512_mask_mov_ps(x, mask, y); + } + static void mask_storeu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_storeu_ps(mem, mask, x); + } + static zmm_t min(zmm_t x, zmm_t y) + { + return _mm512_min_ps(x, y); + } + static zmm_t permutexvar(__m512i idx, zmm_t zmm) + { + return _mm512_permutexvar_ps(idx, zmm); + } + static type_t reducemax(zmm_t v) + { + return _mm512_reduce_max_ps(v); + } + static type_t reducemin(zmm_t v) + { + return _mm512_reduce_min_ps(v); + } + static zmm_t set1(type_t v) + { + return _mm512_set1_ps(v); + } + template <uint8_t mask> + static zmm_t shuffle(zmm_t zmm) + { + return _mm512_shuffle_ps(zmm, zmm, (_MM_PERM_ENUM)mask); + } + static void storeu(void *mem, zmm_t x) + { + return _mm512_storeu_ps(mem, x); + } + + static ymm_t max(ymm_t x, ymm_t y) + { + return _mm256_max_ps(x, y); + } + static ymm_t min(ymm_t x, ymm_t y) + { + return _mm256_min_ps(x, y); + } +}; + +/* + * Assumes zmm is random and performs a full sorting network defined in + * https://en.wikipedia.org/wiki/Bitonic_sorter#/media/File:BitonicSort.svg + */ +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline zmm_t sort_zmm_32bit(zmm_t zmm) +{ + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAA); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(0, 1, 2, 3)>(zmm), + 0xCCCC); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAA); + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_3), zmm), + 0xF0F0); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(1, 0, 3, 2)>(zmm), + 0xCCCC); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAA); + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), zmm), + 0xFF00); + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_6), zmm), + 0xF0F0); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(1, 0, 3, 2)>(zmm), + 0xCCCC); + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAA); + return zmm; +} + +// Assumes zmm is bitonic and performs a recursive half cleaner +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline zmm_t bitonic_merge_zmm_32bit(zmm_t zmm) +{ + // 1) half_cleaner[16]: compare 1-9, 2-10, 3-11 etc .. + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_7), zmm), + 0xFF00); + // 2) half_cleaner[8]: compare 1-5, 2-6, 3-7 etc .. + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_6), zmm), + 0xF0F0); + // 3) half_cleaner[4] + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(1, 0, 3, 2)>(zmm), + 0xCCCC); + // 3) half_cleaner[1] + zmm = cmp_merge<vtype>( + zmm, + vtype::template shuffle<SHUFFLE_MASK(2, 3, 0, 1)>(zmm), + 0xAAAA); + return zmm; +} + +// Assumes zmm1 and zmm2 are sorted and performs a recursive half cleaner +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline void bitonic_merge_two_zmm_32bit(zmm_t *zmm1, zmm_t *zmm2) +{ + // 1) First step of a merging network: coex of zmm1 and zmm2 reversed + *zmm2 = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), *zmm2); + zmm_t zmm3 = vtype::min(*zmm1, *zmm2); + zmm_t zmm4 = vtype::max(*zmm1, *zmm2); + // 2) Recursive half cleaner for each + *zmm1 = bitonic_merge_zmm_32bit<vtype>(zmm3); + *zmm2 = bitonic_merge_zmm_32bit<vtype>(zmm4); +} + +// Assumes [zmm0, zmm1] and [zmm2, zmm3] are sorted and performs a recursive +// half cleaner +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline void bitonic_merge_four_zmm_32bit(zmm_t *zmm) +{ + zmm_t zmm2r = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), zmm[2]); + zmm_t zmm3r = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), zmm[3]); + zmm_t zmm_t1 = vtype::min(zmm[0], zmm3r); + zmm_t zmm_t2 = vtype::min(zmm[1], zmm2r); + zmm_t zmm_t3 = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), + vtype::max(zmm[1], zmm2r)); + zmm_t zmm_t4 = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), + vtype::max(zmm[0], zmm3r)); + zmm_t zmm0 = vtype::min(zmm_t1, zmm_t2); + zmm_t zmm1 = vtype::max(zmm_t1, zmm_t2); + zmm_t zmm2 = vtype::min(zmm_t3, zmm_t4); + zmm_t zmm3 = vtype::max(zmm_t3, zmm_t4); + zmm[0] = bitonic_merge_zmm_32bit<vtype>(zmm0); + zmm[1] = bitonic_merge_zmm_32bit<vtype>(zmm1); + zmm[2] = bitonic_merge_zmm_32bit<vtype>(zmm2); + zmm[3] = bitonic_merge_zmm_32bit<vtype>(zmm3); +} + +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline void bitonic_merge_eight_zmm_32bit(zmm_t *zmm) +{ + zmm_t zmm4r = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), zmm[4]); + zmm_t zmm5r = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), zmm[5]); + zmm_t zmm6r = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), zmm[6]); + zmm_t zmm7r = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), zmm[7]); + zmm_t zmm_t1 = vtype::min(zmm[0], zmm7r); + zmm_t zmm_t2 = vtype::min(zmm[1], zmm6r); + zmm_t zmm_t3 = vtype::min(zmm[2], zmm5r); + zmm_t zmm_t4 = vtype::min(zmm[3], zmm4r); + zmm_t zmm_t5 = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), + vtype::max(zmm[3], zmm4r)); + zmm_t zmm_t6 = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), + vtype::max(zmm[2], zmm5r)); + zmm_t zmm_t7 = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), + vtype::max(zmm[1], zmm6r)); + zmm_t zmm_t8 = vtype::permutexvar(_mm512_set_epi32(NETWORK_32BIT_5), + vtype::max(zmm[0], zmm7r)); + COEX<vtype>(zmm_t1, zmm_t3); + COEX<vtype>(zmm_t2, zmm_t4); + COEX<vtype>(zmm_t5, zmm_t7); + COEX<vtype>(zmm_t6, zmm_t8); + COEX<vtype>(zmm_t1, zmm_t2); + COEX<vtype>(zmm_t3, zmm_t4); + COEX<vtype>(zmm_t5, zmm_t6); + COEX<vtype>(zmm_t7, zmm_t8); + zmm[0] = bitonic_merge_zmm_32bit<vtype>(zmm_t1); + zmm[1] = bitonic_merge_zmm_32bit<vtype>(zmm_t2); + zmm[2] = bitonic_merge_zmm_32bit<vtype>(zmm_t3); + zmm[3] = bitonic_merge_zmm_32bit<vtype>(zmm_t4); + zmm[4] = bitonic_merge_zmm_32bit<vtype>(zmm_t5); + zmm[5] = bitonic_merge_zmm_32bit<vtype>(zmm_t6); + zmm[6] = bitonic_merge_zmm_32bit<vtype>(zmm_t7); + zmm[7] = bitonic_merge_zmm_32bit<vtype>(zmm_t8); +} + +template <typename vtype, typename type_t> +static inline void sort_16_32bit(type_t *arr, int32_t N) +{ + typename vtype::opmask_t load_mask = (0x0001 << N) - 0x0001; + typename vtype::zmm_t zmm + = vtype::mask_loadu(vtype::zmm_max(), load_mask, arr); + vtype::mask_storeu(arr, load_mask, sort_zmm_32bit<vtype>(zmm)); +} + +template <typename vtype, typename type_t> +static inline void sort_32_32bit(type_t *arr, int32_t N) +{ + if (N <= 16) { + sort_16_32bit<vtype>(arr, N); + return; + } + using zmm_t = typename vtype::zmm_t; + zmm_t zmm1 = vtype::loadu(arr); + typename vtype::opmask_t load_mask = (0x0001 << (N - 16)) - 0x0001; + zmm_t zmm2 = vtype::mask_loadu(vtype::zmm_max(), load_mask, arr + 16); + zmm1 = sort_zmm_32bit<vtype>(zmm1); + zmm2 = sort_zmm_32bit<vtype>(zmm2); + bitonic_merge_two_zmm_32bit<vtype>(&zmm1, &zmm2); + vtype::storeu(arr, zmm1); + vtype::mask_storeu(arr + 16, load_mask, zmm2); +} + +template <typename vtype, typename type_t> +static inline void sort_64_32bit(type_t *arr, int32_t N) +{ + if (N <= 32) { + sort_32_32bit<vtype>(arr, N); + return; + } + using zmm_t = typename vtype::zmm_t; + using opmask_t = typename vtype::opmask_t; + zmm_t zmm[4]; + zmm[0] = vtype::loadu(arr); + zmm[1] = vtype::loadu(arr + 16); + opmask_t load_mask1 = 0xFFFF, load_mask2 = 0xFFFF; + uint64_t combined_mask = (0x1ull << (N - 32)) - 0x1ull; + load_mask1 &= combined_mask & 0xFFFF; + load_mask2 &= (combined_mask >> 16) & 0xFFFF; + zmm[2] = vtype::mask_loadu(vtype::zmm_max(), load_mask1, arr + 32); + zmm[3] = vtype::mask_loadu(vtype::zmm_max(), load_mask2, arr + 48); + zmm[0] = sort_zmm_32bit<vtype>(zmm[0]); + zmm[1] = sort_zmm_32bit<vtype>(zmm[1]); + zmm[2] = sort_zmm_32bit<vtype>(zmm[2]); + zmm[3] = sort_zmm_32bit<vtype>(zmm[3]); + bitonic_merge_two_zmm_32bit<vtype>(&zmm[0], &zmm[1]); + bitonic_merge_two_zmm_32bit<vtype>(&zmm[2], &zmm[3]); + bitonic_merge_four_zmm_32bit<vtype>(zmm); + vtype::storeu(arr, zmm[0]); + vtype::storeu(arr + 16, zmm[1]); + vtype::mask_storeu(arr + 32, load_mask1, zmm[2]); + vtype::mask_storeu(arr + 48, load_mask2, zmm[3]); +} + +template <typename vtype, typename type_t> +static inline void sort_128_32bit(type_t *arr, int32_t N) +{ + if (N <= 64) { + sort_64_32bit<vtype>(arr, N); + return; + } + using zmm_t = typename vtype::zmm_t; + using opmask_t = typename vtype::opmask_t; + zmm_t zmm[8]; + zmm[0] = vtype::loadu(arr); + zmm[1] = vtype::loadu(arr + 16); + zmm[2] = vtype::loadu(arr + 32); + zmm[3] = vtype::loadu(arr + 48); + zmm[0] = sort_zmm_32bit<vtype>(zmm[0]); + zmm[1] = sort_zmm_32bit<vtype>(zmm[1]); + zmm[2] = sort_zmm_32bit<vtype>(zmm[2]); + zmm[3] = sort_zmm_32bit<vtype>(zmm[3]); + opmask_t load_mask1 = 0xFFFF, load_mask2 = 0xFFFF; + opmask_t load_mask3 = 0xFFFF, load_mask4 = 0xFFFF; + if (N != 128) { + uint64_t combined_mask = (0x1ull << (N - 64)) - 0x1ull; + load_mask1 &= combined_mask & 0xFFFF; + load_mask2 &= (combined_mask >> 16) & 0xFFFF; + load_mask3 &= (combined_mask >> 32) & 0xFFFF; + load_mask4 &= (combined_mask >> 48) & 0xFFFF; + } + zmm[4] = vtype::mask_loadu(vtype::zmm_max(), load_mask1, arr + 64); + zmm[5] = vtype::mask_loadu(vtype::zmm_max(), load_mask2, arr + 80); + zmm[6] = vtype::mask_loadu(vtype::zmm_max(), load_mask3, arr + 96); + zmm[7] = vtype::mask_loadu(vtype::zmm_max(), load_mask4, arr + 112); + zmm[4] = sort_zmm_32bit<vtype>(zmm[4]); + zmm[5] = sort_zmm_32bit<vtype>(zmm[5]); + zmm[6] = sort_zmm_32bit<vtype>(zmm[6]); + zmm[7] = sort_zmm_32bit<vtype>(zmm[7]); + bitonic_merge_two_zmm_32bit<vtype>(&zmm[0], &zmm[1]); + bitonic_merge_two_zmm_32bit<vtype>(&zmm[2], &zmm[3]); + bitonic_merge_two_zmm_32bit<vtype>(&zmm[4], &zmm[5]); + bitonic_merge_two_zmm_32bit<vtype>(&zmm[6], &zmm[7]); + bitonic_merge_four_zmm_32bit<vtype>(zmm); + bitonic_merge_four_zmm_32bit<vtype>(zmm + 4); + bitonic_merge_eight_zmm_32bit<vtype>(zmm); + vtype::storeu(arr, zmm[0]); + vtype::storeu(arr + 16, zmm[1]); + vtype::storeu(arr + 32, zmm[2]); + vtype::storeu(arr + 48, zmm[3]); + vtype::mask_storeu(arr + 64, load_mask1, zmm[4]); + vtype::mask_storeu(arr + 80, load_mask2, zmm[5]); + vtype::mask_storeu(arr + 96, load_mask3, zmm[6]); + vtype::mask_storeu(arr + 112, load_mask4, zmm[7]); +} + +template <typename vtype, typename type_t> +static inline type_t +get_pivot_32bit(type_t *arr, const int64_t left, const int64_t right) +{ + // median of 16 + int64_t size = (right - left) / 16; + using zmm_t = typename vtype::zmm_t; + using ymm_t = typename vtype::ymm_t; + __m512i rand_index1 = _mm512_set_epi64(left + size, + left + 2 * size, + left + 3 * size, + left + 4 * size, + left + 5 * size, + left + 6 * size, + left + 7 * size, + left + 8 * size); + __m512i rand_index2 = _mm512_set_epi64(left + 9 * size, + left + 10 * size, + left + 11 * size, + left + 12 * size, + left + 13 * size, + left + 14 * size, + left + 15 * size, + left + 16 * size); + ymm_t rand_vec1 + = vtype::template i64gather<sizeof(type_t)>(rand_index1, arr); + ymm_t rand_vec2 + = vtype::template i64gather<sizeof(type_t)>(rand_index2, arr); + zmm_t rand_vec = vtype::merge(rand_vec1, rand_vec2); + zmm_t sort = sort_zmm_32bit<vtype>(rand_vec); + // pivot will never be a nan, since there are no nan's! + return ((type_t *)&sort)[8]; +} + +template <typename vtype, typename type_t> +static inline void +qsort_32bit_(type_t *arr, int64_t left, int64_t right, int64_t max_iters) +{ + /* + * Resort to std::sort if quicksort isnt making any progress + */ + if (max_iters <= 0) { + std::sort(arr + left, arr + right + 1); + return; + } + /* + * Base case: use bitonic networks to sort arrays <= 128 + */ + if (right + 1 - left <= 128) { + sort_128_32bit<vtype>(arr + left, (int32_t)(right + 1 - left)); + return; + } + + type_t pivot = get_pivot_32bit<vtype>(arr, left, right); + type_t smallest = vtype::type_max(); + type_t biggest = vtype::type_min(); + int64_t pivot_index = partition_avx512<vtype>( + arr, left, right + 1, pivot, &smallest, &biggest); + if (pivot != smallest) + qsort_32bit_<vtype>(arr, left, pivot_index - 1, max_iters - 1); + if (pivot != biggest) + qsort_32bit_<vtype>(arr, pivot_index, right, max_iters - 1); +} + +static inline int64_t replace_nan_with_inf(float *arr, int64_t arrsize) +{ + int64_t nan_count = 0; + __mmask16 loadmask = 0xFFFF; + while (arrsize > 0) { + if (arrsize < 16) { loadmask = (0x0001 << arrsize) - 0x0001; } + __m512 in_zmm = _mm512_maskz_loadu_ps(loadmask, arr); + __mmask16 nanmask = _mm512_cmp_ps_mask(in_zmm, in_zmm, _CMP_NEQ_UQ); + nan_count += _mm_popcnt_u32((int32_t)nanmask); + _mm512_mask_storeu_ps(arr, nanmask, ZMM_MAX_FLOAT); + arr += 16; + arrsize -= 16; + } + return nan_count; +} + +static inline void +replace_inf_with_nan(float *arr, int64_t arrsize, int64_t nan_count) +{ + for (int64_t ii = arrsize - 1; nan_count > 0; --ii) { + arr[ii] = std::nanf("1"); + nan_count -= 1; + } +} + +template <> +void avx512_qsort<int32_t>(int32_t *arr, int64_t arrsize) +{ + if (arrsize > 1) { + qsort_32bit_<vector<int32_t>, int32_t>( + arr, 0, arrsize - 1, 2 * (int64_t)log2(arrsize)); + } +} + +template <> +void avx512_qsort<uint32_t>(uint32_t *arr, int64_t arrsize) +{ + if (arrsize > 1) { + qsort_32bit_<vector<uint32_t>, uint32_t>( + arr, 0, arrsize - 1, 2 * (int64_t)log2(arrsize)); + } +} + +template <> +void avx512_qsort<float>(float *arr, int64_t arrsize) +{ + if (arrsize > 1) { + int64_t nan_count = replace_nan_with_inf(arr, arrsize); + qsort_32bit_<vector<float>, float>( + arr, 0, arrsize - 1, 2 * (int64_t)log2(arrsize)); + replace_inf_with_nan(arr, arrsize, nan_count); + } +} + +#endif //__AVX512_QSORT_32BIT__ diff --git a/numpy/core/src/npysort/x86-simd-sort/src/avx512-64bit-qsort.hpp b/numpy/core/src/npysort/x86-simd-sort/src/avx512-64bit-qsort.hpp new file mode 100644 index 000000000..f680c0704 --- /dev/null +++ b/numpy/core/src/npysort/x86-simd-sort/src/avx512-64bit-qsort.hpp @@ -0,0 +1,820 @@ +/******************************************************************* + * Copyright (C) 2022 Intel Corporation + * SPDX-License-Identifier: BSD-3-Clause + * Authors: Raghuveer Devulapalli <raghuveer.devulapalli@intel.com> + * ****************************************************************/ + +#ifndef __AVX512_QSORT_64BIT__ +#define __AVX512_QSORT_64BIT__ + +#include "avx512-common-qsort.h" + +/* + * Constants used in sorting 8 elements in a ZMM registers. Based on Bitonic + * sorting network (see + * https://en.wikipedia.org/wiki/Bitonic_sorter#/media/File:BitonicSort.svg) + */ +// ZMM 7, 6, 5, 4, 3, 2, 1, 0 +#define NETWORK_64BIT_1 4, 5, 6, 7, 0, 1, 2, 3 +#define NETWORK_64BIT_2 0, 1, 2, 3, 4, 5, 6, 7 +#define NETWORK_64BIT_3 5, 4, 7, 6, 1, 0, 3, 2 +#define NETWORK_64BIT_4 3, 2, 1, 0, 7, 6, 5, 4 +static const __m512i rev_index = _mm512_set_epi64(NETWORK_64BIT_2); + +template <> +struct vector<int64_t> { + using type_t = int64_t; + using zmm_t = __m512i; + using ymm_t = __m512i; + using opmask_t = __mmask8; + static const uint8_t numlanes = 8; + + static type_t type_max() + { + return X86_SIMD_SORT_MAX_INT64; + } + static type_t type_min() + { + return X86_SIMD_SORT_MIN_INT64; + } + static zmm_t zmm_max() + { + return _mm512_set1_epi64(type_max()); + } // TODO: this should broadcast bits as is? + + static zmm_t set(type_t v1, + type_t v2, + type_t v3, + type_t v4, + type_t v5, + type_t v6, + type_t v7, + type_t v8) + { + return _mm512_set_epi64(v1, v2, v3, v4, v5, v6, v7, v8); + } + + static opmask_t knot_opmask(opmask_t x) + { + return _knot_mask8(x); + } + static opmask_t ge(zmm_t x, zmm_t y) + { + return _mm512_cmp_epi64_mask(x, y, _MM_CMPINT_NLT); + } + template <int scale> + static zmm_t i64gather(__m512i index, void const *base) + { + return _mm512_i64gather_epi64(index, base, scale); + } + static zmm_t loadu(void const *mem) + { + return _mm512_loadu_si512(mem); + } + static zmm_t max(zmm_t x, zmm_t y) + { + return _mm512_max_epi64(x, y); + } + static void mask_compressstoreu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_compressstoreu_epi64(mem, mask, x); + } + static zmm_t mask_loadu(zmm_t x, opmask_t mask, void const *mem) + { + return _mm512_mask_loadu_epi64(x, mask, mem); + } + static zmm_t mask_mov(zmm_t x, opmask_t mask, zmm_t y) + { + return _mm512_mask_mov_epi64(x, mask, y); + } + static void mask_storeu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_storeu_epi64(mem, mask, x); + } + static zmm_t min(zmm_t x, zmm_t y) + { + return _mm512_min_epi64(x, y); + } + static zmm_t permutexvar(__m512i idx, zmm_t zmm) + { + return _mm512_permutexvar_epi64(idx, zmm); + } + static type_t reducemax(zmm_t v) + { + return _mm512_reduce_max_epi64(v); + } + static type_t reducemin(zmm_t v) + { + return _mm512_reduce_min_epi64(v); + } + static zmm_t set1(type_t v) + { + return _mm512_set1_epi64(v); + } + template <uint8_t mask> + static zmm_t shuffle(zmm_t zmm) + { + __m512d temp = _mm512_castsi512_pd(zmm); + return _mm512_castpd_si512( + _mm512_shuffle_pd(temp, temp, (_MM_PERM_ENUM)mask)); + } + static void storeu(void *mem, zmm_t x) + { + return _mm512_storeu_si512(mem, x); + } +}; +template <> +struct vector<uint64_t> { + using type_t = uint64_t; + using zmm_t = __m512i; + using ymm_t = __m512i; + using opmask_t = __mmask8; + static const uint8_t numlanes = 8; + + static type_t type_max() + { + return X86_SIMD_SORT_MAX_UINT64; + } + static type_t type_min() + { + return 0; + } + static zmm_t zmm_max() + { + return _mm512_set1_epi64(type_max()); + } + + static zmm_t set(type_t v1, + type_t v2, + type_t v3, + type_t v4, + type_t v5, + type_t v6, + type_t v7, + type_t v8) + { + return _mm512_set_epi64(v1, v2, v3, v4, v5, v6, v7, v8); + } + + template <int scale> + static zmm_t i64gather(__m512i index, void const *base) + { + return _mm512_i64gather_epi64(index, base, scale); + } + static opmask_t knot_opmask(opmask_t x) + { + return _knot_mask8(x); + } + static opmask_t ge(zmm_t x, zmm_t y) + { + return _mm512_cmp_epu64_mask(x, y, _MM_CMPINT_NLT); + } + static zmm_t loadu(void const *mem) + { + return _mm512_loadu_si512(mem); + } + static zmm_t max(zmm_t x, zmm_t y) + { + return _mm512_max_epu64(x, y); + } + static void mask_compressstoreu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_compressstoreu_epi64(mem, mask, x); + } + static zmm_t mask_loadu(zmm_t x, opmask_t mask, void const *mem) + { + return _mm512_mask_loadu_epi64(x, mask, mem); + } + static zmm_t mask_mov(zmm_t x, opmask_t mask, zmm_t y) + { + return _mm512_mask_mov_epi64(x, mask, y); + } + static void mask_storeu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_storeu_epi64(mem, mask, x); + } + static zmm_t min(zmm_t x, zmm_t y) + { + return _mm512_min_epu64(x, y); + } + static zmm_t permutexvar(__m512i idx, zmm_t zmm) + { + return _mm512_permutexvar_epi64(idx, zmm); + } + static type_t reducemax(zmm_t v) + { + return _mm512_reduce_max_epu64(v); + } + static type_t reducemin(zmm_t v) + { + return _mm512_reduce_min_epu64(v); + } + static zmm_t set1(type_t v) + { + return _mm512_set1_epi64(v); + } + template <uint8_t mask> + static zmm_t shuffle(zmm_t zmm) + { + __m512d temp = _mm512_castsi512_pd(zmm); + return _mm512_castpd_si512( + _mm512_shuffle_pd(temp, temp, (_MM_PERM_ENUM)mask)); + } + static void storeu(void *mem, zmm_t x) + { + return _mm512_storeu_si512(mem, x); + } +}; +template <> +struct vector<double> { + using type_t = double; + using zmm_t = __m512d; + using ymm_t = __m512d; + using opmask_t = __mmask8; + static const uint8_t numlanes = 8; + + static type_t type_max() + { + return X86_SIMD_SORT_INFINITY; + } + static type_t type_min() + { + return -X86_SIMD_SORT_INFINITY; + } + static zmm_t zmm_max() + { + return _mm512_set1_pd(type_max()); + } + + static zmm_t set(type_t v1, + type_t v2, + type_t v3, + type_t v4, + type_t v5, + type_t v6, + type_t v7, + type_t v8) + { + return _mm512_set_pd(v1, v2, v3, v4, v5, v6, v7, v8); + } + + static opmask_t knot_opmask(opmask_t x) + { + return _knot_mask8(x); + } + static opmask_t ge(zmm_t x, zmm_t y) + { + return _mm512_cmp_pd_mask(x, y, _CMP_GE_OQ); + } + template <int scale> + static zmm_t i64gather(__m512i index, void const *base) + { + return _mm512_i64gather_pd(index, base, scale); + } + static zmm_t loadu(void const *mem) + { + return _mm512_loadu_pd(mem); + } + static zmm_t max(zmm_t x, zmm_t y) + { + return _mm512_max_pd(x, y); + } + static void mask_compressstoreu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_compressstoreu_pd(mem, mask, x); + } + static zmm_t mask_loadu(zmm_t x, opmask_t mask, void const *mem) + { + return _mm512_mask_loadu_pd(x, mask, mem); + } + static zmm_t mask_mov(zmm_t x, opmask_t mask, zmm_t y) + { + return _mm512_mask_mov_pd(x, mask, y); + } + static void mask_storeu(void *mem, opmask_t mask, zmm_t x) + { + return _mm512_mask_storeu_pd(mem, mask, x); + } + static zmm_t min(zmm_t x, zmm_t y) + { + return _mm512_min_pd(x, y); + } + static zmm_t permutexvar(__m512i idx, zmm_t zmm) + { + return _mm512_permutexvar_pd(idx, zmm); + } + static type_t reducemax(zmm_t v) + { + return _mm512_reduce_max_pd(v); + } + static type_t reducemin(zmm_t v) + { + return _mm512_reduce_min_pd(v); + } + static zmm_t set1(type_t v) + { + return _mm512_set1_pd(v); + } + template <uint8_t mask> + static zmm_t shuffle(zmm_t zmm) + { + return _mm512_shuffle_pd(zmm, zmm, (_MM_PERM_ENUM)mask); + } + static void storeu(void *mem, zmm_t x) + { + return _mm512_storeu_pd(mem, x); + } +}; + +/* + * Assumes zmm is random and performs a full sorting network defined in + * https://en.wikipedia.org/wiki/Bitonic_sorter#/media/File:BitonicSort.svg + */ +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline zmm_t sort_zmm_64bit(zmm_t zmm) +{ + zmm = cmp_merge<vtype>( + zmm, vtype::template shuffle<SHUFFLE_MASK(1, 1, 1, 1)>(zmm), 0xAA); + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi64(NETWORK_64BIT_1), zmm), + 0xCC); + zmm = cmp_merge<vtype>( + zmm, vtype::template shuffle<SHUFFLE_MASK(1, 1, 1, 1)>(zmm), 0xAA); + zmm = cmp_merge<vtype>(zmm, vtype::permutexvar(rev_index, zmm), 0xF0); + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi64(NETWORK_64BIT_3), zmm), + 0xCC); + zmm = cmp_merge<vtype>( + zmm, vtype::template shuffle<SHUFFLE_MASK(1, 1, 1, 1)>(zmm), 0xAA); + return zmm; +} + +// Assumes zmm is bitonic and performs a recursive half cleaner +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline zmm_t bitonic_merge_zmm_64bit(zmm_t zmm) +{ + + // 1) half_cleaner[8]: compare 0-4, 1-5, 2-6, 3-7 + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi64(NETWORK_64BIT_4), zmm), + 0xF0); + // 2) half_cleaner[4] + zmm = cmp_merge<vtype>( + zmm, + vtype::permutexvar(_mm512_set_epi64(NETWORK_64BIT_3), zmm), + 0xCC); + // 3) half_cleaner[1] + zmm = cmp_merge<vtype>( + zmm, vtype::template shuffle<SHUFFLE_MASK(1, 1, 1, 1)>(zmm), 0xAA); + return zmm; +} + +// Assumes zmm1 and zmm2 are sorted and performs a recursive half cleaner +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline void bitonic_merge_two_zmm_64bit(zmm_t &zmm1, zmm_t &zmm2) +{ + // 1) First step of a merging network: coex of zmm1 and zmm2 reversed + zmm2 = vtype::permutexvar(rev_index, zmm2); + zmm_t zmm3 = vtype::min(zmm1, zmm2); + zmm_t zmm4 = vtype::max(zmm1, zmm2); + // 2) Recursive half cleaner for each + zmm1 = bitonic_merge_zmm_64bit<vtype>(zmm3); + zmm2 = bitonic_merge_zmm_64bit<vtype>(zmm4); +} + +// Assumes [zmm0, zmm1] and [zmm2, zmm3] are sorted and performs a recursive +// half cleaner +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline void bitonic_merge_four_zmm_64bit(zmm_t *zmm) +{ + // 1) First step of a merging network + zmm_t zmm2r = vtype::permutexvar(rev_index, zmm[2]); + zmm_t zmm3r = vtype::permutexvar(rev_index, zmm[3]); + zmm_t zmm_t1 = vtype::min(zmm[0], zmm3r); + zmm_t zmm_t2 = vtype::min(zmm[1], zmm2r); + // 2) Recursive half clearer: 16 + zmm_t zmm_t3 = vtype::permutexvar(rev_index, vtype::max(zmm[1], zmm2r)); + zmm_t zmm_t4 = vtype::permutexvar(rev_index, vtype::max(zmm[0], zmm3r)); + zmm_t zmm0 = vtype::min(zmm_t1, zmm_t2); + zmm_t zmm1 = vtype::max(zmm_t1, zmm_t2); + zmm_t zmm2 = vtype::min(zmm_t3, zmm_t4); + zmm_t zmm3 = vtype::max(zmm_t3, zmm_t4); + zmm[0] = bitonic_merge_zmm_64bit<vtype>(zmm0); + zmm[1] = bitonic_merge_zmm_64bit<vtype>(zmm1); + zmm[2] = bitonic_merge_zmm_64bit<vtype>(zmm2); + zmm[3] = bitonic_merge_zmm_64bit<vtype>(zmm3); +} + +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline void bitonic_merge_eight_zmm_64bit(zmm_t *zmm) +{ + zmm_t zmm4r = vtype::permutexvar(rev_index, zmm[4]); + zmm_t zmm5r = vtype::permutexvar(rev_index, zmm[5]); + zmm_t zmm6r = vtype::permutexvar(rev_index, zmm[6]); + zmm_t zmm7r = vtype::permutexvar(rev_index, zmm[7]); + zmm_t zmm_t1 = vtype::min(zmm[0], zmm7r); + zmm_t zmm_t2 = vtype::min(zmm[1], zmm6r); + zmm_t zmm_t3 = vtype::min(zmm[2], zmm5r); + zmm_t zmm_t4 = vtype::min(zmm[3], zmm4r); + zmm_t zmm_t5 = vtype::permutexvar(rev_index, vtype::max(zmm[3], zmm4r)); + zmm_t zmm_t6 = vtype::permutexvar(rev_index, vtype::max(zmm[2], zmm5r)); + zmm_t zmm_t7 = vtype::permutexvar(rev_index, vtype::max(zmm[1], zmm6r)); + zmm_t zmm_t8 = vtype::permutexvar(rev_index, vtype::max(zmm[0], zmm7r)); + COEX<vtype>(zmm_t1, zmm_t3); + COEX<vtype>(zmm_t2, zmm_t4); + COEX<vtype>(zmm_t5, zmm_t7); + COEX<vtype>(zmm_t6, zmm_t8); + COEX<vtype>(zmm_t1, zmm_t2); + COEX<vtype>(zmm_t3, zmm_t4); + COEX<vtype>(zmm_t5, zmm_t6); + COEX<vtype>(zmm_t7, zmm_t8); + zmm[0] = bitonic_merge_zmm_64bit<vtype>(zmm_t1); + zmm[1] = bitonic_merge_zmm_64bit<vtype>(zmm_t2); + zmm[2] = bitonic_merge_zmm_64bit<vtype>(zmm_t3); + zmm[3] = bitonic_merge_zmm_64bit<vtype>(zmm_t4); + zmm[4] = bitonic_merge_zmm_64bit<vtype>(zmm_t5); + zmm[5] = bitonic_merge_zmm_64bit<vtype>(zmm_t6); + zmm[6] = bitonic_merge_zmm_64bit<vtype>(zmm_t7); + zmm[7] = bitonic_merge_zmm_64bit<vtype>(zmm_t8); +} + +template <typename vtype, typename zmm_t = typename vtype::zmm_t> +static inline void bitonic_merge_sixteen_zmm_64bit(zmm_t *zmm) +{ + zmm_t zmm8r = vtype::permutexvar(rev_index, zmm[8]); + zmm_t zmm9r = vtype::permutexvar(rev_index, zmm[9]); + zmm_t zmm10r = vtype::permutexvar(rev_index, zmm[10]); + zmm_t zmm11r = vtype::permutexvar(rev_index, zmm[11]); + zmm_t zmm12r = vtype::permutexvar(rev_index, zmm[12]); + zmm_t zmm13r = vtype::permutexvar(rev_index, zmm[13]); + zmm_t zmm14r = vtype::permutexvar(rev_index, zmm[14]); + zmm_t zmm15r = vtype::permutexvar(rev_index, zmm[15]); + zmm_t zmm_t1 = vtype::min(zmm[0], zmm15r); + zmm_t zmm_t2 = vtype::min(zmm[1], zmm14r); + zmm_t zmm_t3 = vtype::min(zmm[2], zmm13r); + zmm_t zmm_t4 = vtype::min(zmm[3], zmm12r); + zmm_t zmm_t5 = vtype::min(zmm[4], zmm11r); + zmm_t zmm_t6 = vtype::min(zmm[5], zmm10r); + zmm_t zmm_t7 = vtype::min(zmm[6], zmm9r); + zmm_t zmm_t8 = vtype::min(zmm[7], zmm8r); + zmm_t zmm_t9 = vtype::permutexvar(rev_index, vtype::max(zmm[7], zmm8r)); + zmm_t zmm_t10 = vtype::permutexvar(rev_index, vtype::max(zmm[6], zmm9r)); + zmm_t zmm_t11 = vtype::permutexvar(rev_index, vtype::max(zmm[5], zmm10r)); + zmm_t zmm_t12 = vtype::permutexvar(rev_index, vtype::max(zmm[4], zmm11r)); + zmm_t zmm_t13 = vtype::permutexvar(rev_index, vtype::max(zmm[3], zmm12r)); + zmm_t zmm_t14 = vtype::permutexvar(rev_index, vtype::max(zmm[2], zmm13r)); + zmm_t zmm_t15 = vtype::permutexvar(rev_index, vtype::max(zmm[1], zmm14r)); + zmm_t zmm_t16 = vtype::permutexvar(rev_index, vtype::max(zmm[0], zmm15r)); + // Recusive half clear 16 zmm regs + COEX<vtype>(zmm_t1, zmm_t5); + COEX<vtype>(zmm_t2, zmm_t6); + COEX<vtype>(zmm_t3, zmm_t7); + COEX<vtype>(zmm_t4, zmm_t8); + COEX<vtype>(zmm_t9, zmm_t13); + COEX<vtype>(zmm_t10, zmm_t14); + COEX<vtype>(zmm_t11, zmm_t15); + COEX<vtype>(zmm_t12, zmm_t16); + // + COEX<vtype>(zmm_t1, zmm_t3); + COEX<vtype>(zmm_t2, zmm_t4); + COEX<vtype>(zmm_t5, zmm_t7); + COEX<vtype>(zmm_t6, zmm_t8); + COEX<vtype>(zmm_t9, zmm_t11); + COEX<vtype>(zmm_t10, zmm_t12); + COEX<vtype>(zmm_t13, zmm_t15); + COEX<vtype>(zmm_t14, zmm_t16); + // + COEX<vtype>(zmm_t1, zmm_t2); + COEX<vtype>(zmm_t3, zmm_t4); + COEX<vtype>(zmm_t5, zmm_t6); + COEX<vtype>(zmm_t7, zmm_t8); + COEX<vtype>(zmm_t9, zmm_t10); + COEX<vtype>(zmm_t11, zmm_t12); + COEX<vtype>(zmm_t13, zmm_t14); + COEX<vtype>(zmm_t15, zmm_t16); + // + zmm[0] = bitonic_merge_zmm_64bit<vtype>(zmm_t1); + zmm[1] = bitonic_merge_zmm_64bit<vtype>(zmm_t2); + zmm[2] = bitonic_merge_zmm_64bit<vtype>(zmm_t3); + zmm[3] = bitonic_merge_zmm_64bit<vtype>(zmm_t4); + zmm[4] = bitonic_merge_zmm_64bit<vtype>(zmm_t5); + zmm[5] = bitonic_merge_zmm_64bit<vtype>(zmm_t6); + zmm[6] = bitonic_merge_zmm_64bit<vtype>(zmm_t7); + zmm[7] = bitonic_merge_zmm_64bit<vtype>(zmm_t8); + zmm[8] = bitonic_merge_zmm_64bit<vtype>(zmm_t9); + zmm[9] = bitonic_merge_zmm_64bit<vtype>(zmm_t10); + zmm[10] = bitonic_merge_zmm_64bit<vtype>(zmm_t11); + zmm[11] = bitonic_merge_zmm_64bit<vtype>(zmm_t12); + zmm[12] = bitonic_merge_zmm_64bit<vtype>(zmm_t13); + zmm[13] = bitonic_merge_zmm_64bit<vtype>(zmm_t14); + zmm[14] = bitonic_merge_zmm_64bit<vtype>(zmm_t15); + zmm[15] = bitonic_merge_zmm_64bit<vtype>(zmm_t16); +} + +template <typename vtype, typename type_t> +static inline void sort_8_64bit(type_t *arr, int32_t N) +{ + typename vtype::opmask_t load_mask = (0x01 << N) - 0x01; + typename vtype::zmm_t zmm + = vtype::mask_loadu(vtype::zmm_max(), load_mask, arr); + vtype::mask_storeu(arr, load_mask, sort_zmm_64bit<vtype>(zmm)); +} + +template <typename vtype, typename type_t> +static inline void sort_16_64bit(type_t *arr, int32_t N) +{ + if (N <= 8) { + sort_8_64bit<vtype>(arr, N); + return; + } + using zmm_t = typename vtype::zmm_t; + zmm_t zmm1 = vtype::loadu(arr); + typename vtype::opmask_t load_mask = (0x01 << (N - 8)) - 0x01; + zmm_t zmm2 = vtype::mask_loadu(vtype::zmm_max(), load_mask, arr + 8); + zmm1 = sort_zmm_64bit<vtype>(zmm1); + zmm2 = sort_zmm_64bit<vtype>(zmm2); + bitonic_merge_two_zmm_64bit<vtype>(zmm1, zmm2); + vtype::storeu(arr, zmm1); + vtype::mask_storeu(arr + 8, load_mask, zmm2); +} + +template <typename vtype, typename type_t> +static inline void sort_32_64bit(type_t *arr, int32_t N) +{ + if (N <= 16) { + sort_16_64bit<vtype>(arr, N); + return; + } + using zmm_t = typename vtype::zmm_t; + using opmask_t = typename vtype::opmask_t; + zmm_t zmm[4]; + zmm[0] = vtype::loadu(arr); + zmm[1] = vtype::loadu(arr + 8); + opmask_t load_mask1 = 0xFF, load_mask2 = 0xFF; + uint64_t combined_mask = (0x1ull << (N - 16)) - 0x1ull; + load_mask1 = (combined_mask)&0xFF; + load_mask2 = (combined_mask >> 8) & 0xFF; + zmm[2] = vtype::mask_loadu(vtype::zmm_max(), load_mask1, arr + 16); + zmm[3] = vtype::mask_loadu(vtype::zmm_max(), load_mask2, arr + 24); + zmm[0] = sort_zmm_64bit<vtype>(zmm[0]); + zmm[1] = sort_zmm_64bit<vtype>(zmm[1]); + zmm[2] = sort_zmm_64bit<vtype>(zmm[2]); + zmm[3] = sort_zmm_64bit<vtype>(zmm[3]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[0], zmm[1]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[2], zmm[3]); + bitonic_merge_four_zmm_64bit<vtype>(zmm); + vtype::storeu(arr, zmm[0]); + vtype::storeu(arr + 8, zmm[1]); + vtype::mask_storeu(arr + 16, load_mask1, zmm[2]); + vtype::mask_storeu(arr + 24, load_mask2, zmm[3]); +} + +template <typename vtype, typename type_t> +static inline void sort_64_64bit(type_t *arr, int32_t N) +{ + if (N <= 32) { + sort_32_64bit<vtype>(arr, N); + return; + } + using zmm_t = typename vtype::zmm_t; + using opmask_t = typename vtype::opmask_t; + zmm_t zmm[8]; + zmm[0] = vtype::loadu(arr); + zmm[1] = vtype::loadu(arr + 8); + zmm[2] = vtype::loadu(arr + 16); + zmm[3] = vtype::loadu(arr + 24); + zmm[0] = sort_zmm_64bit<vtype>(zmm[0]); + zmm[1] = sort_zmm_64bit<vtype>(zmm[1]); + zmm[2] = sort_zmm_64bit<vtype>(zmm[2]); + zmm[3] = sort_zmm_64bit<vtype>(zmm[3]); + opmask_t load_mask1 = 0xFF, load_mask2 = 0xFF; + opmask_t load_mask3 = 0xFF, load_mask4 = 0xFF; + // N-32 >= 1 + uint64_t combined_mask = (0x1ull << (N - 32)) - 0x1ull; + load_mask1 = (combined_mask)&0xFF; + load_mask2 = (combined_mask >> 8) & 0xFF; + load_mask3 = (combined_mask >> 16) & 0xFF; + load_mask4 = (combined_mask >> 24) & 0xFF; + zmm[4] = vtype::mask_loadu(vtype::zmm_max(), load_mask1, arr + 32); + zmm[5] = vtype::mask_loadu(vtype::zmm_max(), load_mask2, arr + 40); + zmm[6] = vtype::mask_loadu(vtype::zmm_max(), load_mask3, arr + 48); + zmm[7] = vtype::mask_loadu(vtype::zmm_max(), load_mask4, arr + 56); + zmm[4] = sort_zmm_64bit<vtype>(zmm[4]); + zmm[5] = sort_zmm_64bit<vtype>(zmm[5]); + zmm[6] = sort_zmm_64bit<vtype>(zmm[6]); + zmm[7] = sort_zmm_64bit<vtype>(zmm[7]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[0], zmm[1]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[2], zmm[3]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[4], zmm[5]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[6], zmm[7]); + bitonic_merge_four_zmm_64bit<vtype>(zmm); + bitonic_merge_four_zmm_64bit<vtype>(zmm + 4); + bitonic_merge_eight_zmm_64bit<vtype>(zmm); + vtype::storeu(arr, zmm[0]); + vtype::storeu(arr + 8, zmm[1]); + vtype::storeu(arr + 16, zmm[2]); + vtype::storeu(arr + 24, zmm[3]); + vtype::mask_storeu(arr + 32, load_mask1, zmm[4]); + vtype::mask_storeu(arr + 40, load_mask2, zmm[5]); + vtype::mask_storeu(arr + 48, load_mask3, zmm[6]); + vtype::mask_storeu(arr + 56, load_mask4, zmm[7]); +} + +template <typename vtype, typename type_t> +static inline void sort_128_64bit(type_t *arr, int32_t N) +{ + if (N <= 64) { + sort_64_64bit<vtype>(arr, N); + return; + } + using zmm_t = typename vtype::zmm_t; + using opmask_t = typename vtype::opmask_t; + zmm_t zmm[16]; + zmm[0] = vtype::loadu(arr); + zmm[1] = vtype::loadu(arr + 8); + zmm[2] = vtype::loadu(arr + 16); + zmm[3] = vtype::loadu(arr + 24); + zmm[4] = vtype::loadu(arr + 32); + zmm[5] = vtype::loadu(arr + 40); + zmm[6] = vtype::loadu(arr + 48); + zmm[7] = vtype::loadu(arr + 56); + zmm[0] = sort_zmm_64bit<vtype>(zmm[0]); + zmm[1] = sort_zmm_64bit<vtype>(zmm[1]); + zmm[2] = sort_zmm_64bit<vtype>(zmm[2]); + zmm[3] = sort_zmm_64bit<vtype>(zmm[3]); + zmm[4] = sort_zmm_64bit<vtype>(zmm[4]); + zmm[5] = sort_zmm_64bit<vtype>(zmm[5]); + zmm[6] = sort_zmm_64bit<vtype>(zmm[6]); + zmm[7] = sort_zmm_64bit<vtype>(zmm[7]); + opmask_t load_mask1 = 0xFF, load_mask2 = 0xFF; + opmask_t load_mask3 = 0xFF, load_mask4 = 0xFF; + opmask_t load_mask5 = 0xFF, load_mask6 = 0xFF; + opmask_t load_mask7 = 0xFF, load_mask8 = 0xFF; + if (N != 128) { + uint64_t combined_mask = (0x1ull << (N - 64)) - 0x1ull; + load_mask1 = (combined_mask)&0xFF; + load_mask2 = (combined_mask >> 8) & 0xFF; + load_mask3 = (combined_mask >> 16) & 0xFF; + load_mask4 = (combined_mask >> 24) & 0xFF; + load_mask5 = (combined_mask >> 32) & 0xFF; + load_mask6 = (combined_mask >> 40) & 0xFF; + load_mask7 = (combined_mask >> 48) & 0xFF; + load_mask8 = (combined_mask >> 56) & 0xFF; + } + zmm[8] = vtype::mask_loadu(vtype::zmm_max(), load_mask1, arr + 64); + zmm[9] = vtype::mask_loadu(vtype::zmm_max(), load_mask2, arr + 72); + zmm[10] = vtype::mask_loadu(vtype::zmm_max(), load_mask3, arr + 80); + zmm[11] = vtype::mask_loadu(vtype::zmm_max(), load_mask4, arr + 88); + zmm[12] = vtype::mask_loadu(vtype::zmm_max(), load_mask5, arr + 96); + zmm[13] = vtype::mask_loadu(vtype::zmm_max(), load_mask6, arr + 104); + zmm[14] = vtype::mask_loadu(vtype::zmm_max(), load_mask7, arr + 112); + zmm[15] = vtype::mask_loadu(vtype::zmm_max(), load_mask8, arr + 120); + zmm[8] = sort_zmm_64bit<vtype>(zmm[8]); + zmm[9] = sort_zmm_64bit<vtype>(zmm[9]); + zmm[10] = sort_zmm_64bit<vtype>(zmm[10]); + zmm[11] = sort_zmm_64bit<vtype>(zmm[11]); + zmm[12] = sort_zmm_64bit<vtype>(zmm[12]); + zmm[13] = sort_zmm_64bit<vtype>(zmm[13]); + zmm[14] = sort_zmm_64bit<vtype>(zmm[14]); + zmm[15] = sort_zmm_64bit<vtype>(zmm[15]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[0], zmm[1]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[2], zmm[3]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[4], zmm[5]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[6], zmm[7]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[8], zmm[9]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[10], zmm[11]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[12], zmm[13]); + bitonic_merge_two_zmm_64bit<vtype>(zmm[14], zmm[15]); + bitonic_merge_four_zmm_64bit<vtype>(zmm); + bitonic_merge_four_zmm_64bit<vtype>(zmm + 4); + bitonic_merge_four_zmm_64bit<vtype>(zmm + 8); + bitonic_merge_four_zmm_64bit<vtype>(zmm + 12); + bitonic_merge_eight_zmm_64bit<vtype>(zmm); + bitonic_merge_eight_zmm_64bit<vtype>(zmm + 8); + bitonic_merge_sixteen_zmm_64bit<vtype>(zmm); + vtype::storeu(arr, zmm[0]); + vtype::storeu(arr + 8, zmm[1]); + vtype::storeu(arr + 16, zmm[2]); + vtype::storeu(arr + 24, zmm[3]); + vtype::storeu(arr + 32, zmm[4]); + vtype::storeu(arr + 40, zmm[5]); + vtype::storeu(arr + 48, zmm[6]); + vtype::storeu(arr + 56, zmm[7]); + vtype::mask_storeu(arr + 64, load_mask1, zmm[8]); + vtype::mask_storeu(arr + 72, load_mask2, zmm[9]); + vtype::mask_storeu(arr + 80, load_mask3, zmm[10]); + vtype::mask_storeu(arr + 88, load_mask4, zmm[11]); + vtype::mask_storeu(arr + 96, load_mask5, zmm[12]); + vtype::mask_storeu(arr + 104, load_mask6, zmm[13]); + vtype::mask_storeu(arr + 112, load_mask7, zmm[14]); + vtype::mask_storeu(arr + 120, load_mask8, zmm[15]); +} + +template <typename vtype, typename type_t> +static inline type_t +get_pivot_64bit(type_t *arr, const int64_t left, const int64_t right) +{ + // median of 8 + int64_t size = (right - left) / 8; + using zmm_t = typename vtype::zmm_t; + __m512i rand_index = _mm512_set_epi64(left + size, + left + 2 * size, + left + 3 * size, + left + 4 * size, + left + 5 * size, + left + 6 * size, + left + 7 * size, + left + 8 * size); + zmm_t rand_vec = vtype::template i64gather<sizeof(type_t)>(rand_index, arr); + // pivot will never be a nan, since there are no nan's! + zmm_t sort = sort_zmm_64bit<vtype>(rand_vec); + return ((type_t *)&sort)[4]; +} + +template <typename vtype, typename type_t> +static inline void +qsort_64bit_(type_t *arr, int64_t left, int64_t right, int64_t max_iters) +{ + /* + * Resort to std::sort if quicksort isnt making any progress + */ + if (max_iters <= 0) { + std::sort(arr + left, arr + right + 1); + return; + } + /* + * Base case: use bitonic networks to sort arrays <= 128 + */ + if (right + 1 - left <= 128) { + sort_128_64bit<vtype>(arr + left, (int32_t)(right + 1 - left)); + return; + } + + type_t pivot = get_pivot_64bit<vtype>(arr, left, right); + type_t smallest = vtype::type_max(); + type_t biggest = vtype::type_min(); + int64_t pivot_index = partition_avx512<vtype>( + arr, left, right + 1, pivot, &smallest, &biggest); + if (pivot != smallest) + qsort_64bit_<vtype>(arr, left, pivot_index - 1, max_iters - 1); + if (pivot != biggest) + qsort_64bit_<vtype>(arr, pivot_index, right, max_iters - 1); +} + +static inline int64_t replace_nan_with_inf(double *arr, int64_t arrsize) +{ + int64_t nan_count = 0; + __mmask8 loadmask = 0xFF; + while (arrsize > 0) { + if (arrsize < 8) { loadmask = (0x01 << arrsize) - 0x01; } + __m512d in_zmm = _mm512_maskz_loadu_pd(loadmask, arr); + __mmask8 nanmask = _mm512_cmp_pd_mask(in_zmm, in_zmm, _CMP_NEQ_UQ); + nan_count += _mm_popcnt_u32((int32_t)nanmask); + _mm512_mask_storeu_pd(arr, nanmask, ZMM_MAX_DOUBLE); + arr += 8; + arrsize -= 8; + } + return nan_count; +} + +static inline void +replace_inf_with_nan(double *arr, int64_t arrsize, int64_t nan_count) +{ + for (int64_t ii = arrsize - 1; nan_count > 0; --ii) { + arr[ii] = std::nan("1"); + nan_count -= 1; + } +} + +template <> +void avx512_qsort<int64_t>(int64_t *arr, int64_t arrsize) +{ + if (arrsize > 1) { + qsort_64bit_<vector<int64_t>, int64_t>( + arr, 0, arrsize - 1, 2 * (int64_t)log2(arrsize)); + } +} + +template <> +void avx512_qsort<uint64_t>(uint64_t *arr, int64_t arrsize) +{ + if (arrsize > 1) { + qsort_64bit_<vector<uint64_t>, uint64_t>( + arr, 0, arrsize - 1, 2 * (int64_t)log2(arrsize)); + } +} + +template <> +void avx512_qsort<double>(double *arr, int64_t arrsize) +{ + if (arrsize > 1) { + int64_t nan_count = replace_nan_with_inf(arr, arrsize); + qsort_64bit_<vector<double>, double>( + arr, 0, arrsize - 1, 2 * (int64_t)log2(arrsize)); + replace_inf_with_nan(arr, arrsize, nan_count); + } +} +#endif // __AVX512_QSORT_64BIT__ diff --git a/numpy/core/src/npysort/x86-simd-sort/src/avx512-common-qsort.h b/numpy/core/src/npysort/x86-simd-sort/src/avx512-common-qsort.h new file mode 100644 index 000000000..e713e1f20 --- /dev/null +++ b/numpy/core/src/npysort/x86-simd-sort/src/avx512-common-qsort.h @@ -0,0 +1,218 @@ +/******************************************************************* + * Copyright (C) 2022 Intel Corporation + * Copyright (C) 2021 Serge Sans Paille + * SPDX-License-Identifier: BSD-3-Clause + * Authors: Raghuveer Devulapalli <raghuveer.devulapalli@intel.com> + * Serge Sans Paille <serge.guelton@telecom-bretagne.eu> + * ****************************************************************/ + +#ifndef __AVX512_QSORT_COMMON__ +#define __AVX512_QSORT_COMMON__ + +/* + * Quicksort using AVX-512. The ideas and code are based on these two research + * papers [1] and [2]. On a high level, the idea is to vectorize quicksort + * partitioning using AVX-512 compressstore instructions. If the array size is + * < 128, then use Bitonic sorting network implemented on 512-bit registers. + * The precise network definitions depend on the dtype and are defined in + * separate files: avx512-16bit-qsort.hpp, avx512-32bit-qsort.hpp and + * avx512-64bit-qsort.hpp. Article [4] is a good resource for bitonic sorting + * network. The core implementations of the vectorized qsort functions + * avx512_qsort<T>(T*, int64_t) are modified versions of avx2 quicksort + * presented in the paper [2] and source code associated with that paper [3]. + * + * [1] Fast and Robust Vectorized In-Place Sorting of Primitive Types + * https://drops.dagstuhl.de/opus/volltexte/2021/13775/ + * + * [2] A Novel Hybrid Quicksort Algorithm Vectorized using AVX-512 on Intel + * Skylake https://arxiv.org/pdf/1704.08579.pdf + * + * [3] https://github.com/simd-sorting/fast-and-robust: SPDX-License-Identifier: MIT + * + * [4] http://mitp-content-server.mit.edu:18180/books/content/sectbyfn?collid=books_pres_0&fn=Chapter%2027.pdf&id=8030 + * + */ + +#include <algorithm> +#include <cmath> +#include <cstdint> +#include <immintrin.h> +#include <limits> + +#define X86_SIMD_SORT_INFINITY std::numeric_limits<double>::infinity() +#define X86_SIMD_SORT_INFINITYF std::numeric_limits<float>::infinity() +#define X86_SIMD_SORT_MAX_UINT16 std::numeric_limits<uint16_t>::max() +#define X86_SIMD_SORT_MAX_INT16 std::numeric_limits<int16_t>::max() +#define X86_SIMD_SORT_MIN_INT16 std::numeric_limits<int16_t>::min() +#define X86_SIMD_SORT_MAX_UINT32 std::numeric_limits<uint32_t>::max() +#define X86_SIMD_SORT_MAX_INT32 std::numeric_limits<int32_t>::max() +#define X86_SIMD_SORT_MIN_INT32 std::numeric_limits<int32_t>::min() +#define X86_SIMD_SORT_MAX_UINT64 std::numeric_limits<uint64_t>::max() +#define X86_SIMD_SORT_MAX_INT64 std::numeric_limits<int64_t>::max() +#define X86_SIMD_SORT_MIN_INT64 std::numeric_limits<int64_t>::min() +#define ZMM_MAX_DOUBLE _mm512_set1_pd(X86_SIMD_SORT_INFINITY) +#define ZMM_MAX_UINT64 _mm512_set1_epi64(X86_SIMD_SORT_MAX_UINT64) +#define ZMM_MAX_INT64 _mm512_set1_epi64(X86_SIMD_SORT_MAX_INT64) +#define ZMM_MAX_FLOAT _mm512_set1_ps(X86_SIMD_SORT_INFINITYF) +#define ZMM_MAX_UINT _mm512_set1_epi32(X86_SIMD_SORT_MAX_UINT32) +#define ZMM_MAX_INT _mm512_set1_epi32(X86_SIMD_SORT_MAX_INT32) +#define ZMM_MAX_UINT16 _mm512_set1_epi16(X86_SIMD_SORT_MAX_UINT16) +#define ZMM_MAX_INT16 _mm512_set1_epi16(X86_SIMD_SORT_MAX_INT16) +#define SHUFFLE_MASK(a, b, c, d) (a << 6) | (b << 4) | (c << 2) | d + +template <typename type> +struct vector; + +template <typename T> +void avx512_qsort(T *arr, int64_t arrsize); + +/* + * COEX == Compare and Exchange two registers by swapping min and max values + */ +template <typename vtype, typename mm_t> +static void COEX(mm_t &a, mm_t &b) +{ + mm_t temp = a; + a = vtype::min(a, b); + b = vtype::max(temp, b); +} + +template <typename vtype, + typename zmm_t = typename vtype::zmm_t, + typename opmask_t = typename vtype::opmask_t> +static inline zmm_t cmp_merge(zmm_t in1, zmm_t in2, opmask_t mask) +{ + zmm_t min = vtype::min(in2, in1); + zmm_t max = vtype::max(in2, in1); + return vtype::mask_mov(min, mask, max); // 0 -> min, 1 -> max +} + +/* + * Parition one ZMM register based on the pivot and returns the index of the + * last element that is less than equal to the pivot. + */ +template <typename vtype, typename type_t, typename zmm_t> +static inline int32_t partition_vec(type_t *arr, + int64_t left, + int64_t right, + const zmm_t curr_vec, + const zmm_t pivot_vec, + zmm_t *smallest_vec, + zmm_t *biggest_vec) +{ + /* which elements are larger than the pivot */ + typename vtype::opmask_t gt_mask = vtype::ge(curr_vec, pivot_vec); + int32_t amount_gt_pivot = _mm_popcnt_u32((int32_t)gt_mask); + vtype::mask_compressstoreu( + arr + left, vtype::knot_opmask(gt_mask), curr_vec); + vtype::mask_compressstoreu( + arr + right - amount_gt_pivot, gt_mask, curr_vec); + *smallest_vec = vtype::min(curr_vec, *smallest_vec); + *biggest_vec = vtype::max(curr_vec, *biggest_vec); + return amount_gt_pivot; +} + +/* + * Parition an array based on the pivot and returns the index of the + * last element that is less than equal to the pivot. + */ +template <typename vtype, typename type_t> +static inline int64_t partition_avx512(type_t *arr, + int64_t left, + int64_t right, + type_t pivot, + type_t *smallest, + type_t *biggest) +{ + /* make array length divisible by vtype::numlanes , shortening the array */ + for (int32_t i = (right - left) % vtype::numlanes; i > 0; --i) { + *smallest = std::min(*smallest, arr[left]); + *biggest = std::max(*biggest, arr[left]); + if (arr[left] > pivot) { std::swap(arr[left], arr[--right]); } + else { + ++left; + } + } + + if (left == right) + return left; /* less than vtype::numlanes elements in the array */ + + using zmm_t = typename vtype::zmm_t; + zmm_t pivot_vec = vtype::set1(pivot); + zmm_t min_vec = vtype::set1(*smallest); + zmm_t max_vec = vtype::set1(*biggest); + + if (right - left == vtype::numlanes) { + zmm_t vec = vtype::loadu(arr + left); + int32_t amount_gt_pivot = partition_vec<vtype>(arr, + left, + left + vtype::numlanes, + vec, + pivot_vec, + &min_vec, + &max_vec); + *smallest = vtype::reducemin(min_vec); + *biggest = vtype::reducemax(max_vec); + return left + (vtype::numlanes - amount_gt_pivot); + } + + // first and last vtype::numlanes values are partitioned at the end + zmm_t vec_left = vtype::loadu(arr + left); + zmm_t vec_right = vtype::loadu(arr + (right - vtype::numlanes)); + // store points of the vectors + int64_t r_store = right - vtype::numlanes; + int64_t l_store = left; + // indices for loading the elements + left += vtype::numlanes; + right -= vtype::numlanes; + while (right - left != 0) { + zmm_t curr_vec; + /* + * if fewer elements are stored on the right side of the array, + * then next elements are loaded from the right side, + * otherwise from the left side + */ + if ((r_store + vtype::numlanes) - right < left - l_store) { + right -= vtype::numlanes; + curr_vec = vtype::loadu(arr + right); + } + else { + curr_vec = vtype::loadu(arr + left); + left += vtype::numlanes; + } + // partition the current vector and save it on both sides of the array + int32_t amount_gt_pivot + = partition_vec<vtype>(arr, + l_store, + r_store + vtype::numlanes, + curr_vec, + pivot_vec, + &min_vec, + &max_vec); + ; + r_store -= amount_gt_pivot; + l_store += (vtype::numlanes - amount_gt_pivot); + } + + /* partition and save vec_left and vec_right */ + int32_t amount_gt_pivot = partition_vec<vtype>(arr, + l_store, + r_store + vtype::numlanes, + vec_left, + pivot_vec, + &min_vec, + &max_vec); + l_store += (vtype::numlanes - amount_gt_pivot); + amount_gt_pivot = partition_vec<vtype>(arr, + l_store, + l_store + vtype::numlanes, + vec_right, + pivot_vec, + &min_vec, + &max_vec); + l_store += (vtype::numlanes - amount_gt_pivot); + *smallest = vtype::reducemin(min_vec); + *biggest = vtype::reducemax(max_vec); + return l_store; +} +#endif // __AVX512_QSORT_COMMON__ |
