summaryrefslogtreecommitdiff
path: root/numpy/core/src/common
diff options
context:
space:
mode:
authorQiyu8 <fangchunlin@huawei.com>2021-01-25 10:42:05 +0800
committerQiyu8 <fangchunlin@huawei.com>2021-01-25 10:42:05 +0800
commitd244aa9bae95d6061feaec4a9873ef4992c26245 (patch)
treefb2fd6d48cb92fac2d3a1090fcadab30b79d3110 /numpy/core/src/common
parent9fa688a9c433aad96c9d53c7fea09d54efbe5b68 (diff)
downloadnumpy-d244aa9bae95d6061feaec4a9873ef4992c26245.tar.gz
improve sumup intriniscs.
Diffstat (limited to 'numpy/core/src/common')
-rw-r--r--numpy/core/src/common/simd/avx2/arithmetic.h41
-rw-r--r--numpy/core/src/common/simd/avx512/arithmetic.h36
-rw-r--r--numpy/core/src/common/simd/neon/arithmetic.h42
-rw-r--r--numpy/core/src/common/simd/sse/arithmetic.h40
-rw-r--r--numpy/core/src/common/simd/sse/utils.h4
-rw-r--r--numpy/core/src/common/simd/vsx/arithmetic.h41
6 files changed, 111 insertions, 93 deletions
diff --git a/numpy/core/src/common/simd/avx2/arithmetic.h b/numpy/core/src/common/simd/avx2/arithmetic.h
index 9e13d6324..c4c5f2093 100644
--- a/numpy/core/src/common/simd/avx2/arithmetic.h
+++ b/numpy/core/src/common/simd/avx2/arithmetic.h
@@ -118,16 +118,10 @@
}
#endif // !NPY_HAVE_FMA3
-// Horizontal add: Calculates the sum of all vector elements.
-
-NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
-{
- __m256i four = _mm256_sad_epu8(a, _mm256_setzero_si256());
- __m128i two = _mm_add_epi16(_mm256_castsi256_si128(four), _mm256_extracti128_si256(four, 1));
- __m128i one = _mm_add_epi16(two, _mm_unpackhi_epi64(two, two));
- return (npy_uint16)_mm_cvtsi128_si32(one);
-}
-
+/***************************
+ * Summation
+ ***************************/
+// reduce sum across vector
NPY_FINLINE npy_uint32 npyv_sum_u32(npyv_u32 a)
{
__m256i s0 = _mm256_hadd_epi32(a, a);
@@ -137,15 +131,6 @@ NPY_FINLINE npy_uint32 npyv_sum_u32(npyv_u32 a)
return _mm_cvtsi128_si32(s1);
}
-NPY_FINLINE npy_uint32 npyv_sumup_u16(npyv_u16 a)
-{
- const npyv_u16 even_mask = _mm256_set1_epi32(0x0000FFFF);
- __m256i even = _mm256_and_si256(a, even_mask);
- __m256i odd = _mm256_srli_epi32(a, 16);
- __m256i eight = _mm256_add_epi32(even, odd);
- return npyv_sum_u32(eight);
-}
-
NPY_FINLINE npy_uint64 npyv_sum_u64(npyv_u64 a)
{
__m256i two = _mm256_add_epi64(a, _mm256_shuffle_epi32(a, _MM_SHUFFLE(1, 0, 3, 2)));
@@ -172,6 +157,24 @@ NPY_FINLINE double npyv_sum_f64(npyv_f64 a)
return _mm_cvtsd_f64(sum);
}
+// extend sum across vector
+NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
+{
+ __m256i four = _mm256_sad_epu8(a, _mm256_setzero_si256());
+ __m128i two = _mm_add_epi16(_mm256_castsi256_si128(four), _mm256_extracti128_si256(four, 1));
+ __m128i one = _mm_add_epi16(two, _mm_unpackhi_epi64(two, two));
+ return (npy_uint16)_mm_cvtsi128_si32(one);
+}
+
+NPY_FINLINE npy_uint32 npyv_sumup_u16(npyv_u16 a)
+{
+ const npyv_u16 even_mask = _mm256_set1_epi32(0x0000FFFF);
+ __m256i even = _mm256_and_si256(a, even_mask);
+ __m256i odd = _mm256_srli_epi32(a, 16);
+ __m256i eight = _mm256_add_epi32(even, odd);
+ return npyv_sum_u32(eight);
+}
+
#endif // _NPY_SIMD_AVX2_ARITHMETIC_H
diff --git a/numpy/core/src/common/simd/avx512/arithmetic.h b/numpy/core/src/common/simd/avx512/arithmetic.h
index ea7dc0c3c..a6e448bae 100644
--- a/numpy/core/src/common/simd/avx512/arithmetic.h
+++ b/numpy/core/src/common/simd/avx512/arithmetic.h
@@ -130,7 +130,7 @@ NPY_FINLINE __m512i npyv_mul_u8(__m512i a, __m512i b)
#define npyv_nmulsub_f64 _mm512_fnmsub_pd
/***************************
- * Reduce Sum: Calculates the sum of all vector elements.
+ * Summation: Calculates the sum of all vector elements.
* there are three ways to implement reduce sum for AVX512:
* 1- split(256) /add /split(128) /add /hadd /hadd /extract
* 2- shuff(cross) /add /shuff(cross) /add /shuff /add /shuff /add /extract
@@ -144,29 +144,13 @@ NPY_FINLINE __m512i npyv_mul_u8(__m512i a, __m512i b)
* The third one is almost the same as the second one but only works for
* intel compiler/GCC 7.1/Clang 4, we still need to support older GCC.
***************************/
-
-NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
-{
-#ifdef NPY_HAVE_AVX512BW
- __m512i eight = _mm512_sad_epu8(a, _mm512_setzero_si512());
- __m256i four = _mm256_add_epi16(npyv512_lower_si256(eight), npyv512_higher_si256(eight));
-#else
- __m256i lo_four = _mm256_sad_epu8(npyv512_lower_si256(a), _mm256_setzero_si256());
- __m256i hi_four = _mm256_sad_epu8(npyv512_higher_si256(a), _mm256_setzero_si256());
- __m256i four = _mm256_add_epi16(lo_four, hi_four);
-#endif
- __m128i two = _mm_add_epi16(_mm256_castsi256_si128(four), _mm256_extracti128_si256(four, 1));
- __m128i one = _mm_add_epi16(two, _mm_unpackhi_epi64(two, two));
- return (npy_uint16)_mm_cvtsi128_si32(one);
-}
-
+// reduce sum across vector
#ifdef NPY_HAVE_AVX512F_REDUCE
#define npyv_sum_u32 _mm512_reduce_add_epi32
#define npyv_sum_u64 _mm512_reduce_add_epi64
#define npyv_sum_f32 _mm512_reduce_add_ps
#define npyv_sum_f64 _mm512_reduce_add_pd
#else
-
NPY_FINLINE npy_uint32 npyv_sum_u32(npyv_u32 a)
{
__m256i half = _mm256_add_epi32(npyv512_lower_si256(a), npyv512_higher_si256(a));
@@ -208,6 +192,22 @@ NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
}
#endif
+// extend sum across vector
+NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
+{
+#ifdef NPY_HAVE_AVX512BW
+ __m512i eight = _mm512_sad_epu8(a, _mm512_setzero_si512());
+ __m256i four = _mm256_add_epi16(npyv512_lower_si256(eight), npyv512_higher_si256(eight));
+#else
+ __m256i lo_four = _mm256_sad_epu8(npyv512_lower_si256(a), _mm256_setzero_si256());
+ __m256i hi_four = _mm256_sad_epu8(npyv512_higher_si256(a), _mm256_setzero_si256());
+ __m256i four = _mm256_add_epi16(lo_four, hi_four);
+#endif
+ __m128i two = _mm_add_epi16(_mm256_castsi256_si128(four), _mm256_extracti128_si256(four, 1));
+ __m128i one = _mm_add_epi16(two, _mm_unpackhi_epi64(two, two));
+ return (npy_uint16)_mm_cvtsi128_si32(one);
+}
+
NPY_FINLINE npy_uint32 npyv_sumup_u16(npyv_u16 a)
{
const npyv_u16 even_mask = _mm512_set1_epi32(0x0000FFFF);
diff --git a/numpy/core/src/common/simd/neon/arithmetic.h b/numpy/core/src/common/simd/neon/arithmetic.h
index 81207ea5e..af34299a0 100644
--- a/numpy/core/src/common/simd/neon/arithmetic.h
+++ b/numpy/core/src/common/simd/neon/arithmetic.h
@@ -131,30 +131,16 @@
{ return vfmsq_f64(vnegq_f64(c), a, b); }
#endif // NPY_SIMD_F64
-// Horizontal add: Calculates the sum of all vector elements.
+/***************************
+ * Summation
+ ***************************/
+// reduce sum across vector
#if NPY_SIMD_F64
- #define npyv_sumup_u8 vaddlvq_u8
- #define npyv_sumup_u16 vaddlvq_u16
#define npyv_sum_u32 vaddvq_u32
#define npyv_sum_u64 vaddvq_u64
#define npyv_sum_f32 vaddvq_f32
#define npyv_sum_f64 vaddvq_f64
#else
-
- NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
- {
- uint32x4_t t0 = vpaddlq_u16(vpaddlq_u8(a));
- uint32x2_t t1 = vpadd_u32(vget_low_u32(t0), vget_high_u32(t0));
- return vget_lane_u32(vpadd_u32(t1, t1), 0);
- }
-
- NPY_FINLINE npy_uint32 npyv_sumup_u16(npyv_u16 a)
- {
- uint32x4_t t0 = vpaddlq_u16(a);
- uint32x2_t t1 = vpadd_u32(vget_low_u32(t0), vget_high_u32(t0));
- return vget_lane_u32(vpadd_u32(t1, t1), 0);
- }
-
NPY_FINLINE npy_uint64 npyv_sum_u64(npyv_u64 a)
{
return vget_lane_u64(vadd_u64(vget_low_u64(a), vget_high_u64(a)),0);
@@ -173,4 +159,24 @@
}
#endif
+// extend sum across vector
+#if NPY_SIMD_F64
+ #define npyv_sumup_u8 vaddlvq_u8
+ #define npyv_sumup_u16 vaddlvq_u16
+#else
+ NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
+ {
+ uint32x4_t t0 = vpaddlq_u16(vpaddlq_u8(a));
+ uint32x2_t t1 = vpadd_u32(vget_low_u32(t0), vget_high_u32(t0));
+ return vget_lane_u32(vpadd_u32(t1, t1), 0);
+ }
+
+ NPY_FINLINE npy_uint32 npyv_sumup_u16(npyv_u16 a)
+ {
+ uint32x4_t t0 = vpaddlq_u16(a);
+ uint32x2_t t1 = vpadd_u32(vget_low_u32(t0), vget_high_u32(t0));
+ return vget_lane_u32(vpadd_u32(t1, t1), 0);
+ }
+#endif
+
#endif // _NPY_SIMD_NEON_ARITHMETIC_H
diff --git a/numpy/core/src/common/simd/sse/arithmetic.h b/numpy/core/src/common/simd/sse/arithmetic.h
index 92a53e630..fcb0a1716 100644
--- a/numpy/core/src/common/simd/sse/arithmetic.h
+++ b/numpy/core/src/common/simd/sse/arithmetic.h
@@ -148,14 +148,10 @@ NPY_FINLINE __m128i npyv_mul_u8(__m128i a, __m128i b)
}
#endif // !NPY_HAVE_FMA3
-// Horizontal add: Calculates the sum of all vector elements.
-
-NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
-{
- __m128i half = _mm_sad_epu8(a, _mm_setzero_si128());
- return (unsigned)_mm_cvtsi128_si32(_mm_add_epi32(half, _mm_unpackhi_epi64(half, half)));
-}
-
+/***************************
+ * Summation
+ ***************************/
+// reduce sum across vector
NPY_FINLINE npy_uint32 npyv_sum_u32(npyv_u32 a)
{
__m128i t = _mm_add_epi32(a, _mm_srli_si128(a, 8));
@@ -163,17 +159,10 @@ NPY_FINLINE npy_uint32 npyv_sum_u32(npyv_u32 a)
return (unsigned)_mm_cvtsi128_si32(t);
}
-NPY_FINLINE npy_uint32 npyv_sumup_u16(npyv_u16 a)
-{
- npyv_u32x2 res = npyv_expand_u32_u16(a);
- return (unsigned)npyv_sum_u32(_mm_add_epi32(res.val[0], res.val[1]));
-}
-
NPY_FINLINE npy_uint64 npyv_sum_u64(npyv_u64 a)
{
- npy_uint64 NPY_DECL_ALIGNED(32) idx[2];
- npyv_storea_u64(idx, a);
- return idx[0] + idx[1];
+ __m128i one = _mm_add_epi64(a, _mm_unpackhi_epi64(a, a));
+ return (npy_uint64)npyv128_cvtsi128_si64(one);
}
NPY_FINLINE float npyv_sum_f32(npyv_f32 a)
@@ -199,6 +188,23 @@ NPY_FINLINE double npyv_sum_f64(npyv_f64 a)
#endif
}
+// extend sum across vector
+NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
+{
+ __m128i two = _mm_sad_epu8(a, _mm_setzero_si128());
+ __m128i one = _mm_add_epi16(two, _mm_unpackhi_epi64(two, two));
+ return (npy_uint16)_mm_cvtsi128_si32(one);
+}
+
+NPY_FINLINE npy_uint32 npyv_sumup_u16(npyv_u16 a)
+{
+ const __m128i even_mask = _mm_set1_epi32(0x0000FFFF);
+ __m128i even = _mm_and_si128(a, even_mask);
+ __m128i odd = _mm_srli_epi32(a, 16);
+ __m128i four = _mm_add_epi32(even, odd);
+ return npyv_sum_u32(four);
+}
+
#endif // _NPY_SIMD_SSE_ARITHMETIC_H
diff --git a/numpy/core/src/common/simd/sse/utils.h b/numpy/core/src/common/simd/sse/utils.h
index 5e03e12a3..c23def11d 100644
--- a/numpy/core/src/common/simd/sse/utils.h
+++ b/numpy/core/src/common/simd/sse/utils.h
@@ -6,9 +6,9 @@
#define _NPY_SIMD_SSE_UTILS_H
#if !defined(__x86_64__) && !defined(_M_X64)
-NPY_FINLINE npy_uint64 npyv128_cvtsi128_si64(__m128i a)
+NPY_FINLINE npy_int64 npyv128_cvtsi128_si64(__m128i a)
{
- npy_uint64 NPY_DECL_ALIGNED(32) idx[2];
+ npy_int64 NPY_DECL_ALIGNED(16) idx[2];
_mm_store_si128((__m128i *)idx, a);
return idx[0];
}
diff --git a/numpy/core/src/common/simd/vsx/arithmetic.h b/numpy/core/src/common/simd/vsx/arithmetic.h
index 97d5efe61..339677857 100644
--- a/numpy/core/src/common/simd/vsx/arithmetic.h
+++ b/numpy/core/src/common/simd/vsx/arithmetic.h
@@ -116,25 +116,10 @@
#define npyv_nmulsub_f32 vec_nmadd // equivalent to -(a*b + c)
#define npyv_nmulsub_f64 vec_nmadd
-// Horizontal add: Calculates the sum of all vector elements.
-
-NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
-{
- const npyv_u32 zero = npyv_zero_u32();
- npyv_u32 four = vec_sum4s(a, zero);
- npyv_u32 one = vec_sums((npyv_s32)four, (npyv_s32)zero);
- return (npy_uint16)vec_extract(one, 3);
-}
-
-NPY_FINLINE npy_uint32 npyv_sumup_u16(npyv_u16 a)
-{
- const npyv_s32 zero = npyv_zero_s32();
- npyv_u32x2 eight = npyv_expand_u32_u16(a);
- npyv_u32 four = vec_add(eight.val[0], eight.val[1]);
- npyv_s32 one = vec_sums((npyv_s32)four, zero);
- return (npy_uint32)vec_extract(one, 3);
-}
-
+/***************************
+ * Summation
+ ***************************/
+// reduce sum across vector
NPY_FINLINE npy_uint64 npyv_sum_u64(npyv_u64 a)
{
return vec_extract(vec_add(a, vec_mergel(a, a)), 0);
@@ -157,4 +142,22 @@ NPY_FINLINE double npyv_sum_f64(npyv_f64 a)
return vec_extract(a, 0) + vec_extract(a, 1);
}
+// extend sum across vector
+NPY_FINLINE npy_uint16 npyv_sumup_u8(npyv_u8 a)
+{
+ const npyv_u32 zero = npyv_zero_u32();
+ npyv_u32 four = vec_sum4s(a, zero);
+ npyv_s32 one = vec_sums((npyv_s32)four, (npyv_s32)zero);
+ return (npy_uint16)vec_extract(one, 3);
+}
+
+NPY_FINLINE npy_uint32 npyv_sumup_u16(npyv_u16 a)
+{
+ const npyv_s32 zero = npyv_zero_s32();
+ npyv_u32x2 eight = npyv_expand_u32_u16(a);
+ npyv_u32 four = vec_add(eight.val[0], eight.val[1]);
+ npyv_s32 one = vec_sums((npyv_s32)four, zero);
+ return (npy_uint32)vec_extract(one, 3);
+}
+
#endif // _NPY_SIMD_VSX_ARITHMETIC_H