--- a/include/everett/select_groups.h +++ b/include/everett/select_groups.h @@ -24 +24 @@ -#if defined(__BMI2__) +#if defined(__BMI2__) || defined(__AVX2__) @@ -25,0 +26,3 @@ +#endif +#if defined(__aarch64__) && defined(__ARM_NEON) +#include @@ -41,0 +45,21 @@ + } + + inline std::array populations4(std::uint64_t const * words) noexcept { +#if defined(__aarch64__) && defined(__ARM_NEON) + auto a = vpaddlq_u32(vpaddlq_u16(vpaddlq_u8(vcntq_u8(vld1q_u8(reinterpret_cast(words)))))); + auto b = vpaddlq_u32(vpaddlq_u16(vpaddlq_u8(vcntq_u8(vld1q_u8(reinterpret_cast(words + 2)))))); + return {unsigned(vgetq_lane_u64(a, 0)), unsigned(vgetq_lane_u64(a, 1)), + unsigned(vgetq_lane_u64(b, 0)), unsigned(vgetq_lane_u64(b, 1))}; +#elif defined(__AVX2__) + auto table = _mm256_setr_epi8(0,1,1,2,1,2,2,3,1,2,2,3,2,3,3,4,0,1,1,2,1,2,2,3,1,2,2,3,2,3,3,4); + auto bits = _mm256_loadu_si256(reinterpret_cast<__m256i const *>(words)); + auto mask = _mm256_set1_epi8(15); + auto bytes = _mm256_add_epi8(_mm256_shuffle_epi8(table, _mm256_and_si256(bits, mask)), + _mm256_shuffle_epi8(table, _mm256_and_si256(_mm256_srli_epi16(bits, 4), mask))); + auto counts = _mm256_sad_epu8(bytes, _mm256_setzero_si256()); + return {unsigned(_mm256_extract_epi64(counts, 0)), unsigned(_mm256_extract_epi64(counts, 1)), + unsigned(_mm256_extract_epi64(counts, 2)), unsigned(_mm256_extract_epi64(counts, 3))}; +#else + return {unsigned(std::popcount(words[0])), unsigned(std::popcount(words[1])), + unsigned(std::popcount(words[2])), unsigned(std::popcount(words[3]))}; +#endif @@ -107,3 +131,20 @@ - for (; source.size() - at >= inputs; at += inputs) - low_tile(source.data() + at, out.data() + (at / inputs) * outputs, - std::make_index_sequence{}); + for (; source.size() - at >= inputs; at += inputs) { + auto destination = out.data() + (at / inputs) * outputs; +#if defined(__aarch64__) && defined(__ARM_NEON) && __BYTE_ORDER__ == __ORDER_LITTLE_ENDIAN__ + if constexpr (W == 32) { + vst1_u32(reinterpret_cast(destination), vmovn_u64(vld1q_u64(source.data() + at))); + } else if constexpr (W == 16) { + auto a = vmovn_u64(vld1q_u64(source.data() + at)); + auto b = vmovn_u64(vld1q_u64(source.data() + at + 2)); + vst1_u16(reinterpret_cast(destination), vmovn_u32(vcombine_u32(a, b))); + } else if constexpr (W == 8) { + auto a = vmovn_u64(vld1q_u64(source.data() + at)); + auto b = vmovn_u64(vld1q_u64(source.data() + at + 2)); + auto c = vmovn_u64(vld1q_u64(source.data() + at + 4)); + auto d = vmovn_u64(vld1q_u64(source.data() + at + 6)); + vst1_u8(reinterpret_cast(destination), + vmovn_u16(vcombine_u16(vmovn_u32(vcombine_u32(a, b)), vmovn_u32(vcombine_u32(c, d))))); + } else +#endif + low_tile(source.data() + at, destination, std::make_index_sequence{}); + } @@ -254,0 +296,10 @@ + if (scanned && high_.size() - word >= 4 && scanned <= 61) { + auto populations = select_groups_detail::populations4(high_.data() + word); + auto total = populations[0] + populations[1] + populations[2] + populations[3]; + if (remaining >= total) { remaining -= total; scanned += 3; word += 3; continue; } + unsigned lane = 0; + while (remaining >= populations[lane]) { remaining -= populations[lane]; ++lane; } + word += lane; + value = high_[word]; + population = populations[lane]; + } @@ -286,2 +337,9 @@ - for (std::size_t i = 1; i < residuals.size(); ++i) - if (residuals[i] < residuals[i - 1]) + std::size_t checked = 1; +#if defined(__aarch64__) && defined(__ARM_NEON) + for (; residuals.size() - checked >= 2; checked += 2) + if (vmaxvq_u32(vreinterpretq_u32_u64(vcltq_u64(vld1q_u64(residuals.data() + checked), + vld1q_u64(residuals.data() + checked - 1))))) + throw std::invalid_argument("select_groups nonmonotone offsets"); +#endif + for (; checked < residuals.size(); ++checked) + if (residuals[checked] < residuals[checked - 1])