diff --git a/include/xsimd/arch/common/xsimd_common_arithmetic.hpp b/include/xsimd/arch/common/xsimd_common_arithmetic.hpp index 0c4d2c437..e6eedab4a 100644 --- a/include/xsimd/arch/common/xsimd_common_arithmetic.hpp +++ b/include/xsimd/arch/common/xsimd_common_arithmetic.hpp @@ -330,33 +330,74 @@ namespace xsimd } // rotl + // Rotations go through the unsigned type: an arithmetic right shift of a + // negative value would fill the vacated bits with the sign instead of + // the bits shifted out on the left. template XSIMD_INLINE batch rotl(batch const& self, STy other, requires_arch) noexcept { + using U = std::make_unsigned_t; constexpr auto bits = std::numeric_limits::digits + std::numeric_limits::is_signed; - return (self << other) | (self >> (bits - other)); + auto const u = bitwise_cast(self); + if constexpr (std::is_integral_v) + { + return bitwise_cast((u << other) | (u >> (bits - other))); + } + else + { + auto const o = bitwise_cast(other); + return bitwise_cast((u << o) | (u >> (batch(bits) - o))); + } } template XSIMD_INLINE batch rotl(batch const& self, requires_arch) noexcept { + using U = std::make_unsigned_t; constexpr auto bits = std::numeric_limits::digits + std::numeric_limits::is_signed; static_assert(count < bits, "Count amount must be less than the number of bits in T"); - return bitwise_lshift(self) | bitwise_rshift(self); + if constexpr (count == 0) + { + return self; + } + else + { + auto const u = bitwise_cast(self); + return bitwise_cast(bitwise_lshift(u) | bitwise_rshift(u)); + } } // rotr template XSIMD_INLINE batch rotr(batch const& self, STy other, requires_arch) noexcept { + using U = std::make_unsigned_t; constexpr auto bits = std::numeric_limits::digits + std::numeric_limits::is_signed; - return (self >> other) | (self << (bits - other)); + auto const u = bitwise_cast(self); + if constexpr (std::is_integral_v) + { + return bitwise_cast((u >> other) | (u << (bits - other))); + } + else + { + auto const o = bitwise_cast(other); + return bitwise_cast((u >> o) | (u << (batch(bits) - o))); + } } template XSIMD_INLINE batch rotr(batch const& self, requires_arch) noexcept { + using U = std::make_unsigned_t; constexpr auto bits = std::numeric_limits::digits + std::numeric_limits::is_signed; static_assert(count < bits, "Count must be less than the number of bits in T"); - return bitwise_rshift(self) | bitwise_lshift(self); + if constexpr (count == 0) + { + return self; + } + else + { + auto const u = bitwise_cast(self); + return bitwise_cast(bitwise_rshift(u) | bitwise_lshift(u)); + } } // sadd diff --git a/include/xsimd/arch/xsimd_avx2.hpp b/include/xsimd/arch/xsimd_avx2.hpp index ba6825cb8..4d16def95 100644 --- a/include/xsimd/arch/xsimd_avx2.hpp +++ b/include/xsimd/arch/xsimd_avx2.hpp @@ -422,10 +422,9 @@ namespace xsimd { if constexpr (sizeof(T) == 1) { - // 8-bit left shift via 16-bit shift + mask + // 8-bit right shift via 16-bit shift + mask of the bits that stay in the byte const __m256i shifted = _mm256_srli_epi16(self, shift); - // TODO(C++17): without `if constexpr ` we must ensure the compile-time shift does not overflow - constexpr uint8_t mask8 = static_cast(sizeof(T) == 1 ? ((1u << shift) - 1u) : 0); + constexpr uint8_t mask8 = static_cast(0xFFu >> shift); const __m256i mask = _mm256_set1_epi8(mask8); return _mm256_and_si256(shifted, mask); } diff --git a/include/xsimd/arch/xsimd_scalar.hpp b/include/xsimd/arch/xsimd_scalar.hpp index 6adc9dce5..b822e0638 100644 --- a/include/xsimd/arch/xsimd_scalar.hpp +++ b/include/xsimd/arch/xsimd_scalar.hpp @@ -418,12 +418,18 @@ namespace xsimd return 1. / x; } + // Rotations go through the unsigned type: an arithmetic right shift of a + // negative value would fill the vacated bits with the sign instead of the + // bits shifted out on the left. template XSIMD_INLINE std::enable_if_t && std::is_integral_v, T0> rotl(T0 x, T1 shift) noexcept { - constexpr auto bits = std::numeric_limits::digits + std::numeric_limits::is_signed; - return (x << shift) | (x >> (bits - shift)); + using U = std::make_unsigned_t; + constexpr unsigned bits = sizeof(T0) * 8; + auto const u = static_cast(x); + auto const s = static_cast(shift) & (bits - 1); + return static_cast(static_cast(u << s) | static_cast(u >> ((bits - s) & (bits - 1)))); } template XSIMD_INLINE std::enable_if_t, T> @@ -431,15 +437,18 @@ namespace xsimd { constexpr auto bits = std::numeric_limits::digits + std::numeric_limits::is_signed; static_assert(count < bits, "Count must be less than the number of bits in T"); - return (x << count) | (x >> (bits - count)); + return rotl(x, count); } template XSIMD_INLINE std::enable_if_t && std::is_integral_v, T0> rotr(T0 x, T1 shift) noexcept { - constexpr auto bits = std::numeric_limits::digits + std::numeric_limits::is_signed; - return (x >> shift) | (x << (bits - shift)); + using U = std::make_unsigned_t; + constexpr unsigned bits = sizeof(T0) * 8; + auto const u = static_cast(x); + auto const s = static_cast(shift) & (bits - 1); + return static_cast(static_cast(u >> s) | static_cast(u << ((bits - s) & (bits - 1)))); } template XSIMD_INLINE std::enable_if_t, T> @@ -447,7 +456,7 @@ namespace xsimd { constexpr auto bits = std::numeric_limits::digits + std::numeric_limits::is_signed; static_assert(count < bits, "Count must be less than the number of bits in T"); - return (x >> count) | (x << (bits - count)); + return rotr(x, count); } template diff --git a/include/xsimd/arch/xsimd_sse2.hpp b/include/xsimd/arch/xsimd_sse2.hpp index 58e35254d..26e858d9c 100644 --- a/include/xsimd/arch/xsimd_sse2.hpp +++ b/include/xsimd/arch/xsimd_sse2.hpp @@ -500,12 +500,13 @@ namespace xsimd { if constexpr (sizeof(T) == 1) { - // 8-bit arithmetic right shift via 16-bit shift + sign-extension handling. - __m128i shifted = _mm_srai_epi16(self, static_cast(shift)); - __m128i sign_mask = _mm_set1_epi16(static_cast(0xFF00 >> shift)); - __m128i cmp_negative = _mm_cmpgt_epi8(_mm_setzero_si128(), self); - return _mm_or_si128(_mm_and_si128(sign_mask, cmp_negative), - _mm_andnot_si128(sign_mask, shifted)); + // 8-bit arithmetic right shift: logical shift on 16-bit lanes, drop the + // bits that came from the neighbouring byte, then sign-extend. + constexpr uint8_t keep = static_cast(0xFFu >> shift); + constexpr uint8_t sign = static_cast(0x80u >> shift); + __m128i shifted = _mm_and_si128(_mm_srli_epi16(self, static_cast(shift)), _mm_set1_epi8(static_cast(keep))); + __m128i sign_bit = _mm_set1_epi8(static_cast(sign)); + return _mm_sub_epi8(_mm_xor_si128(shifted, sign_bit), sign_bit); } else if constexpr (sizeof(T) == 2) { @@ -522,10 +523,9 @@ namespace xsimd { if constexpr (sizeof(T) == 1) { - // 8-bit left shift via 16-bit shift + mask + // 8-bit right shift via 16-bit shift + mask of the bits that stay in the byte __m128i shifted = _mm_srli_epi16(self, static_cast(shift)); - // TODO(C++17): without `if constexpr ` we must ensure the compile-time shift does not overflow - constexpr uint8_t mask8 = static_cast(sizeof(T) == 1 ? ((1u << shift) - 1u) : 0); + constexpr uint8_t mask8 = static_cast(0xFFu >> shift); const __m128i mask = _mm_set1_epi8(mask8); return _mm_and_si128(shifted, mask); } diff --git a/test/test_xsimd_api.cpp b/test/test_xsimd_api.cpp index bc8d8f6bd..fa62367d3 100644 --- a/test/test_xsimd_api.cpp +++ b/test/test_xsimd_api.cpp @@ -14,6 +14,10 @@ #include +#include +#include +#include + template struct scalar_type { @@ -519,6 +523,67 @@ TEST_CASE_TEMPLATE("[xsimd api | ssub at type minimum]", B, INTEGRAL_TYPES) } } +// Rotating a negative value must reintroduce the bits shifted out, not copies +// of the sign bit. The expected values are computed on the unsigned type. +TEST_CASE_TEMPLATE("[xsimd api | rotations of the sign bit]", B, INTEGRAL_TYPES) +{ + using value_type = typename scalar_type::type; + using U = std::make_unsigned_t; + constexpr int bits = sizeof(value_type) * 8; + auto ref_rotl = [](U u, int n) + { return static_cast(static_cast(static_cast(u << n) | static_cast(u >> (bits - n)))); }; + auto ref_rotr = [](U u, int n) + { return static_cast(static_cast(static_cast(u >> n) | static_cast(u << (bits - n)))); }; + + // 1 followed by zeros, and a pattern with the sign bit and the low bit set + U const inputs[] = { static_cast(U(1) << (bits - 1)), static_cast((U(1) << (bits - 1)) | U(1)) }; + for (U u : inputs) + { + value_type const v = static_cast(u); + CHECK_EQ(extract(xsimd::rotl(B(v), B(value_type(1)))), ref_rotl(u, 1)); + CHECK_EQ(extract(xsimd::rotl(B(v), 3)), ref_rotl(u, 3)); + CHECK_EQ(extract(xsimd::rotl<1>(B(v))), ref_rotl(u, 1)); + CHECK_EQ(extract(xsimd::rotl(B(v))), ref_rotl(u, bits - 1)); + CHECK_EQ(extract(xsimd::rotr(B(v), B(value_type(1)))), ref_rotr(u, 1)); + CHECK_EQ(extract(xsimd::rotr(B(v), 3)), ref_rotr(u, 3)); + CHECK_EQ(extract(xsimd::rotr<1>(B(v))), ref_rotr(u, 1)); + CHECK_EQ(extract(xsimd::rotr(B(v))), ref_rotr(u, bits - 1)); + } +} + +// The 8-bit fixed right shift is built from a 16-bit shift and a mask that +// removes the bits coming from the neighbouring byte. +template +void check_fixed_rshift_8bit(std::index_sequence) +{ + using B = xsimd::batch; + std::array in, out; + for (int base = 0; base < 256; ++base) + { + for (size_t i = 0; i < B::size; ++i) + { + in[i] = static_cast(base + 37 * i); + } + B const v = B::load_unaligned(in.data()); + auto check = [&](auto shift_tag) + { + constexpr size_t shift = decltype(shift_tag)::value; + xsimd::bitwise_rshift(v).store_unaligned(out.data()); + for (size_t i = 0; i < B::size; ++i) + { + CHECK_EQ(out[i], static_cast(in[i] >> shift)); + } + }; + (check(std::integral_constant {}), ...); + } +} + +TEST_CASE("[xsimd api | fixed right shift of 8-bit integers]") +{ + check_fixed_rshift_8bit(std::make_index_sequence<8> {}); + check_fixed_rshift_8bit(std::make_index_sequence<8> {}); +} + /* * Functions that apply on floating points types only */