summaryrefslogtreecommitdiff
path: root/numpy/core/src
diff options
context:
space:
mode:
authorRaghuveer Devulapalli <raghuveer.devulapalli@intel.com>2022-09-19 10:35:03 -0700
committerRaghuveer Devulapalli <raghuveer.devulapalli@intel.com>2023-01-30 13:38:39 -0800
commit49278b961b7254bc6a4aee478587c69682a3827e (patch)
tree9bcf5e3f022df96018d2dcd158578c16218b829f /numpy/core/src
parentc662a712a30b1b640a80421619bb97556ffe965b (diff)
downloadnumpy-49278b961b7254bc6a4aee478587c69682a3827e.tar.gz
ENH: Add x86-simd-sort source files
Diffstat (limited to 'numpy/core/src')
-rw-r--r--numpy/core/src/npysort/x86-simd-sort/src/avx512-16bit-qsort.hpp527
-rw-r--r--numpy/core/src/npysort/x86-simd-sort/src/avx512-32bit-qsort.hpp712
-rw-r--r--numpy/core/src/npysort/x86-simd-sort/src/avx512-64bit-qsort.hpp820
-rw-r--r--numpy/core/src/npysort/x86-simd-sort/src/avx512-common-qsort.h218
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__