Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 45 additions & 4 deletions include/xsimd/arch/common/xsimd_common_arithmetic.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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 <class A, class T, class STy>
XSIMD_INLINE batch<T, A> rotl(batch<T, A> const& self, STy other, requires_arch<common>) noexcept
{
using U = std::make_unsigned_t<T>;
constexpr auto bits = std::numeric_limits<T>::digits + std::numeric_limits<T>::is_signed;
return (self << other) | (self >> (bits - other));
auto const u = bitwise_cast<U>(self);
if constexpr (std::is_integral_v<STy>)
{
return bitwise_cast<T>((u << other) | (u >> (bits - other)));
}
else
{
auto const o = bitwise_cast<U>(other);
return bitwise_cast<T>((u << o) | (u >> (batch<U, A>(bits) - o)));
}
}
template <size_t count, class A, class T>
XSIMD_INLINE batch<T, A> rotl(batch<T, A> const& self, requires_arch<common>) noexcept
{
using U = std::make_unsigned_t<T>;
constexpr auto bits = std::numeric_limits<T>::digits + std::numeric_limits<T>::is_signed;
static_assert(count < bits, "Count amount must be less than the number of bits in T");
return bitwise_lshift<count>(self) | bitwise_rshift<bits - count>(self);
if constexpr (count == 0)
{
return self;
}
else
{
auto const u = bitwise_cast<U>(self);
return bitwise_cast<T>(bitwise_lshift<count>(u) | bitwise_rshift<bits - count>(u));
}
}

// rotr
template <class A, class T, class STy>
XSIMD_INLINE batch<T, A> rotr(batch<T, A> const& self, STy other, requires_arch<common>) noexcept
{
using U = std::make_unsigned_t<T>;
constexpr auto bits = std::numeric_limits<T>::digits + std::numeric_limits<T>::is_signed;
return (self >> other) | (self << (bits - other));
auto const u = bitwise_cast<U>(self);
if constexpr (std::is_integral_v<STy>)
{
return bitwise_cast<T>((u >> other) | (u << (bits - other)));
}
else
{
auto const o = bitwise_cast<U>(other);
return bitwise_cast<T>((u >> o) | (u << (batch<U, A>(bits) - o)));
}
}
template <size_t count, class A, class T>
XSIMD_INLINE batch<T, A> rotr(batch<T, A> const& self, requires_arch<common>) noexcept
{
using U = std::make_unsigned_t<T>;
constexpr auto bits = std::numeric_limits<T>::digits + std::numeric_limits<T>::is_signed;
static_assert(count < bits, "Count must be less than the number of bits in T");
return bitwise_rshift<count>(self) | bitwise_lshift<bits - count>(self);
if constexpr (count == 0)
{
return self;
}
else
{
auto const u = bitwise_cast<U>(self);
return bitwise_cast<T>(bitwise_rshift<count>(u) | bitwise_lshift<bits - count>(u));
}
}

// sadd
Expand Down
5 changes: 2 additions & 3 deletions include/xsimd/arch/xsimd_avx2.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<uint8_t>(sizeof(T) == 1 ? ((1u << shift) - 1u) : 0);
constexpr uint8_t mask8 = static_cast<uint8_t>(0xFFu >> shift);
const __m256i mask = _mm256_set1_epi8(mask8);
return _mm256_and_si256(shifted, mask);
}
Expand Down
21 changes: 15 additions & 6 deletions include/xsimd/arch/xsimd_scalar.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -418,36 +418,45 @@ 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 <class T0, class T1>
XSIMD_INLINE std::enable_if_t<std::is_integral_v<T0> && std::is_integral_v<T1>, T0>
rotl(T0 x, T1 shift) noexcept
{
constexpr auto bits = std::numeric_limits<T0>::digits + std::numeric_limits<T0>::is_signed;
return (x << shift) | (x >> (bits - shift));
using U = std::make_unsigned_t<T0>;
constexpr unsigned bits = sizeof(T0) * 8;
auto const u = static_cast<U>(x);
auto const s = static_cast<unsigned>(shift) & (bits - 1);
return static_cast<T0>(static_cast<U>(u << s) | static_cast<U>(u >> ((bits - s) & (bits - 1))));
}
template <size_t count, class T>
XSIMD_INLINE std::enable_if_t<std::is_integral_v<T>, T>
rotl(T x) noexcept
{
constexpr auto bits = std::numeric_limits<T>::digits + std::numeric_limits<T>::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 <class T0, class T1>
XSIMD_INLINE std::enable_if_t<std::is_integral_v<T0> && std::is_integral_v<T1>, T0>
rotr(T0 x, T1 shift) noexcept
{
constexpr auto bits = std::numeric_limits<T0>::digits + std::numeric_limits<T0>::is_signed;
return (x >> shift) | (x << (bits - shift));
using U = std::make_unsigned_t<T0>;
constexpr unsigned bits = sizeof(T0) * 8;
auto const u = static_cast<U>(x);
auto const s = static_cast<unsigned>(shift) & (bits - 1);
return static_cast<T0>(static_cast<U>(u >> s) | static_cast<U>(u << ((bits - s) & (bits - 1))));
}
template <size_t count, class T>
XSIMD_INLINE std::enable_if_t<std::is_integral_v<T>, T>
rotr(T x) noexcept
{
constexpr auto bits = std::numeric_limits<T>::digits + std::numeric_limits<T>::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 <class T>
Expand Down
18 changes: 9 additions & 9 deletions include/xsimd/arch/xsimd_sse2.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -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<int>(shift));
__m128i sign_mask = _mm_set1_epi16(static_cast<short>(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<uint8_t>(0xFFu >> shift);
constexpr uint8_t sign = static_cast<uint8_t>(0x80u >> shift);
__m128i shifted = _mm_and_si128(_mm_srli_epi16(self, static_cast<int>(shift)), _mm_set1_epi8(static_cast<char>(keep)));
__m128i sign_bit = _mm_set1_epi8(static_cast<char>(sign));
return _mm_sub_epi8(_mm_xor_si128(shifted, sign_bit), sign_bit);
}
else if constexpr (sizeof(T) == 2)
{
Expand All @@ -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<int>(shift));
// TODO(C++17): without `if constexpr ` we must ensure the compile-time shift does not overflow
constexpr uint8_t mask8 = static_cast<uint8_t>(sizeof(T) == 1 ? ((1u << shift) - 1u) : 0);
constexpr uint8_t mask8 = static_cast<uint8_t>(0xFFu >> shift);
const __m128i mask = _mm_set1_epi8(mask8);
return _mm_and_si128(shifted, mask);
}
Expand Down
65 changes: 65 additions & 0 deletions test/test_xsimd_api.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,10 @@

#include <doctest/doctest.h>

#include <array>
#include <type_traits>
#include <utility>

template <class T>
struct scalar_type
{
Expand Down Expand Up @@ -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<B>::type;
using U = std::make_unsigned_t<value_type>;
constexpr int bits = sizeof(value_type) * 8;
auto ref_rotl = [](U u, int n)
{ return static_cast<value_type>(static_cast<U>(static_cast<U>(u << n) | static_cast<U>(u >> (bits - n)))); };
auto ref_rotr = [](U u, int n)
{ return static_cast<value_type>(static_cast<U>(static_cast<U>(u >> n) | static_cast<U>(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>(U(1) << (bits - 1)), static_cast<U>((U(1) << (bits - 1)) | U(1)) };
for (U u : inputs)
{
value_type const v = static_cast<value_type>(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<bits - 1>(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<bits - 1>(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 <class T, size_t... Shifts>
void check_fixed_rshift_8bit(std::index_sequence<Shifts...>)
{
using B = xsimd::batch<T>;
std::array<T, B::size> in, out;
for (int base = 0; base < 256; ++base)
{
for (size_t i = 0; i < B::size; ++i)
{
in[i] = static_cast<T>(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<shift>(v).store_unaligned(out.data());
for (size_t i = 0; i < B::size; ++i)
{
CHECK_EQ(out[i], static_cast<T>(in[i] >> shift));
}
};
(check(std::integral_constant<size_t, Shifts> {}), ...);
}
}

TEST_CASE("[xsimd api | fixed right shift of 8-bit integers]")
{
check_fixed_rshift_8bit<uint8_t>(std::make_index_sequence<8> {});
check_fixed_rshift_8bit<int8_t>(std::make_index_sequence<8> {});
}

/*
* Functions that apply on floating points types only
*/
Expand Down
Loading