Ship the TLS mesh, composer, Beaver/Yao/leaf MPC, prep/online paths, apps, and docs so the tree is pushable before elevating share_expr, security_mode, and prep resume. Co-authored-by: Cursor <cursoragent@cursor.com>
1379 lines
47 KiB
C++
1379 lines
47 KiB
C++
/// @file dpf/leaf_arithmetic.hpp
|
||
/// @brief Addition, subtraction, and multiplication of packed leaves.
|
||
/// @author Ryan Henry <ryan.henry@ucalgary.ca>
|
||
/// @copyright Copyright (c) 2019-2024 Ryan Henry and [others](@ref authors)
|
||
/// @license Released under a GNU General Public v2.0 (GPLv2) license;
|
||
/// see [LICENSE.md](@ref license) for details.
|
||
|
||
#ifndef LIBDPF_INCLUDE_DPF_LEAF_ARITHMETIC_HPP__
|
||
#define LIBDPF_INCLUDE_DPF_LEAF_ARITHMETIC_HPP__
|
||
|
||
#include "hedley/hedley.h"
|
||
|
||
#include <cstddef>
|
||
#include <cstring>
|
||
#include <type_traits>
|
||
#include <functional>
|
||
#include <algorithm>
|
||
#include <iterator>
|
||
#include <array>
|
||
|
||
#include "simde/simde/x86/avx2.h"
|
||
#include "portable-snippets/exact-int/exact-int.h"
|
||
|
||
#include "dpf/bit.hpp"
|
||
#include "dpf/bitstring.hpp"
|
||
#include "dpf/blob.hpp"
|
||
#include "dpf/packed_lane_arithmetic.hpp"
|
||
#include "dpf/wildcard.hpp"
|
||
#include "dpf/xor_wrapper.hpp"
|
||
|
||
namespace dpf
|
||
{
|
||
|
||
namespace leaf_arithmetic
|
||
{
|
||
|
||
template <typename OutputT, typename NodeT, typename Enable = void> struct add_t;
|
||
template <typename OutputT, typename NodeT, typename Enable = void> struct subtract_t;
|
||
template <typename OutputT, typename NodeT, typename Enable = void> struct multiply_t;
|
||
|
||
template <typename OutputT>
|
||
struct add_t<OutputT, void>
|
||
{
|
||
template <typename NodeT>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const NodeT & a, const NodeT & b) const noexcept
|
||
{
|
||
static constexpr auto adder = add_t<dpf::concrete_type_t<OutputT>, NodeT>{};
|
||
return adder(a, b);
|
||
}
|
||
};
|
||
|
||
template <typename OutputT>
|
||
struct subtract_t<OutputT, void>
|
||
{
|
||
|
||
template <typename NodeT>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const NodeT & a, const NodeT & b) const noexcept
|
||
{
|
||
static constexpr auto subtracter = subtract_t<dpf::concrete_type_t<OutputT>, NodeT>{};
|
||
return subtracter(a, b);
|
||
}
|
||
};
|
||
|
||
template <>
|
||
struct multiply_t<void, void>
|
||
{
|
||
template <typename NodeT, typename OutputT>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const NodeT & a, OutputT b) const noexcept
|
||
{
|
||
static constexpr auto multiplier = multiply_t<OutputT, NodeT>{};
|
||
return multiplier(a, b);
|
||
}
|
||
};
|
||
|
||
} // namespace leaf_arithmetic
|
||
|
||
template <typename OutputT>
|
||
static constexpr auto add_leaf = leaf_arithmetic::add_t<OutputT, void>{};
|
||
|
||
template <typename OutputT>
|
||
static constexpr auto subtract_leaf = leaf_arithmetic::subtract_t<OutputT, void>{};
|
||
|
||
static constexpr auto multiply_leaf = leaf_arithmetic::multiply_t<void, void>{};
|
||
|
||
/// @brief Scalar add in the leaf output group.
|
||
/// @details `float`/`double` use XOR of the IEEE bit pattern (same as
|
||
/// `add_t` on packed leaves), not IEEE floating-point addition.
|
||
/// \complexity O(1) for a scalar. A packed node walks its bytes once (`add_t` / the lane loops): O(sizeof node) and the carry stays inside a `twobit` or `nyble` lane.
|
||
template <typename T>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
T leaf_group_add(const T & a, const T & b) noexcept
|
||
{
|
||
if constexpr (std::is_same_v<T, float> || std::is_same_v<T, double>)
|
||
{
|
||
using bits_t = std::conditional_t<sizeof(T) == 4, psnip_uint32_t, psnip_uint64_t>;
|
||
bits_t aa{}, bb{};
|
||
utils::raw_memcpy(&aa, &a, sizeof(T));
|
||
utils::raw_memcpy(&bb, &b, sizeof(T));
|
||
bits_t cc = static_cast<bits_t>(aa ^ bb);
|
||
T out{};
|
||
utils::raw_memcpy(&out, &cc, sizeof(T));
|
||
return out;
|
||
}
|
||
else
|
||
{
|
||
return static_cast<T>(a + b);
|
||
}
|
||
}
|
||
|
||
/// @brief Scalar multiply in the leaf output group (`float`/`double`: AND of bits).
|
||
template <typename T>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
T leaf_group_mul(const T & a, const T & b) noexcept
|
||
{
|
||
if constexpr (std::is_same_v<T, float> || std::is_same_v<T, double>)
|
||
{
|
||
using bits_t = std::conditional_t<sizeof(T) == 4, psnip_uint32_t, psnip_uint64_t>;
|
||
bits_t aa{}, bb{};
|
||
utils::raw_memcpy(&aa, &a, sizeof(T));
|
||
utils::raw_memcpy(&bb, &b, sizeof(T));
|
||
bits_t cc = static_cast<bits_t>(aa & bb);
|
||
T out{};
|
||
utils::raw_memcpy(&out, &cc, sizeof(T));
|
||
return out;
|
||
}
|
||
else
|
||
{
|
||
return static_cast<T>(a * b);
|
||
}
|
||
}
|
||
|
||
/// @brief Multiplicative identity for leaf scaling (all-ones for XOR/AND groups).
|
||
template <typename T>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
T leaf_group_one() noexcept
|
||
{
|
||
if constexpr (std::is_same_v<T, float> || std::is_same_v<T, double>)
|
||
{
|
||
using bits_t = std::conditional_t<sizeof(T) == 4, psnip_uint32_t, psnip_uint64_t>;
|
||
bits_t ones = static_cast<bits_t>(~bits_t{0});
|
||
T out{};
|
||
utils::raw_memcpy(&out, &ones, sizeof(T));
|
||
return out;
|
||
}
|
||
else if constexpr (utils::is_xor_wrapper_v<T>)
|
||
{
|
||
using u = typename T::value_type;
|
||
return T{static_cast<u>(~u{0})};
|
||
}
|
||
else
|
||
{
|
||
return T{1};
|
||
}
|
||
}
|
||
|
||
namespace leaf_arithmetic
|
||
{
|
||
|
||
namespace detail
|
||
{
|
||
|
||
template <std::size_t Nbits,
|
||
typename WordT>
|
||
struct bitstring_xor_t
|
||
{
|
||
template <typename NodeT>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const NodeT & a, const NodeT & b) const noexcept
|
||
{
|
||
return a ^ b;
|
||
}
|
||
|
||
template <typename NodeT,
|
||
std::size_t N>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const std::array<NodeT, N> & a, const std::array<NodeT, N> & b) const noexcept
|
||
{
|
||
std::array<NodeT, N> ret;
|
||
std::transform(std::begin(a), std::end(a), std::begin(b), std::begin(ret), [](const NodeT & a_, const NodeT & b_)
|
||
{
|
||
return a_ ^ b_;
|
||
});
|
||
return ret;
|
||
}
|
||
};
|
||
|
||
// /// @brief adds vectors of 8-bit integral types
|
||
// template <typename NodeT> struct add8_t;
|
||
// /// @brief adds vectors of 16-bit integral types
|
||
// template <typename NodeT> struct add16_t;
|
||
// /// @brief adds vectors of 32-bit integral types
|
||
// template <typename NodeT> struct add32_t;
|
||
// /// @brief adds vectors of 64-bit integral types
|
||
// template <typename NodeT> struct add64_t;
|
||
|
||
/// @brief Function object for adding vectors of `16x8`-bit integral types
|
||
struct add16x8_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return simde_mm_add_epi8(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for adding vectors of `8x16`-bit integral types
|
||
struct add8x16_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return simde_mm_add_epi16(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for adding vectors of `4x32`-bit integral types
|
||
struct add4x32_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return simde_mm_add_epi32(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for adding vectors of `2x64`-bit integral types
|
||
struct add2x64_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return simde_mm_add_epi64(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for adding vectors of `32x8`-bit integral types
|
||
struct add32x8_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return simde_mm256_add_epi8(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for adding vectors of `16x16`-bit integral types
|
||
struct add16x16_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return simde_mm256_add_epi16(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for adding vectors of `8x32`-bit integral types
|
||
struct add8x32_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return simde_mm256_add_epi32(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for adding vectors of `4x64`-bit integral types
|
||
struct add4x64_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return simde_mm256_add_epi64(a, b);
|
||
}
|
||
};
|
||
|
||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||
template <typename OutputT, typename Enabled = void>
|
||
struct add_array_t
|
||
{
|
||
template <typename T, std::size_t N>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const std::array<T, N> & a, const std::array<T, N> & b) const noexcept
|
||
{
|
||
std::array<T, N> c;
|
||
std::transform(std::begin(a), std::end(a), std::begin(b), std::begin(c),
|
||
[](const T & a, const T & b) { return std::bit_xor<>{}(a, b); });
|
||
return c;
|
||
}
|
||
};
|
||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||
|
||
template <typename OutputT>
|
||
struct add_array_t<OutputT, std::enable_if_t<dpf::utils::has_operators_plus_minus_v<OutputT>, void>>
|
||
{
|
||
using output_type = OutputT;
|
||
template <typename T, std::size_t N>
|
||
auto operator()(const std::array<T, N> & a, const std::array<T, N> & b) const
|
||
{
|
||
static_assert(sizeof(output_type) == sizeof(std::array<T, N>),
|
||
"arithmetic leaf array and output type must be the same size");
|
||
std::array<T, N> c;
|
||
output_type a_, b_;
|
||
utils::raw_memcpy(&a_, std::data(a), sizeof(a_));
|
||
utils::raw_memcpy(&b_, std::data(b), sizeof(b_));
|
||
output_type c_ = a_ + b_;
|
||
utils::raw_memcpy(std::data(c), &c_, sizeof(c_));
|
||
return c;
|
||
}
|
||
};
|
||
|
||
} // namespace detail
|
||
|
||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||
|
||
template <> struct add_t<bool, simde__m128i> final : public detail::add16x8_t {};
|
||
template <> struct add_t<char, simde__m128i> final : public detail::add16x8_t {};
|
||
// template <> struct add_t<unsigned char, simde__m128i> final : public detail::add16x8_t {};
|
||
template <> struct add_t<psnip_int8_t, simde__m128i> final : public detail::add16x8_t {};
|
||
template <> struct add_t<psnip_uint8_t, simde__m128i> final : public detail::add16x8_t {};
|
||
|
||
template <> struct add_t<psnip_int16_t, simde__m128i> final : public detail::add8x16_t {};
|
||
template <> struct add_t<psnip_uint16_t, simde__m128i> final : public detail::add8x16_t {};
|
||
|
||
template <> struct add_t<psnip_int32_t, simde__m128i> final : public detail::add4x32_t {};
|
||
template <> struct add_t<psnip_uint32_t, simde__m128i> final : public detail::add4x32_t {};
|
||
|
||
template <> struct add_t<psnip_int64_t, simde__m128i> final : public detail::add2x64_t {};
|
||
template <> struct add_t<psnip_uint64_t, simde__m128i> final : public detail::add2x64_t {};
|
||
|
||
template <> struct add_t<simde_int128, simde__m128i> final
|
||
{
|
||
auto operator()(const simde__m128i & lhs, const simde__m128i & rhs) const
|
||
{
|
||
simde__m128i ret;
|
||
simde_int128 lhs_, rhs_;
|
||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_int128));
|
||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_int128));
|
||
simde_int128 sum = lhs_ + rhs_;
|
||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m128i));
|
||
return ret;
|
||
}
|
||
};
|
||
template <> struct add_t<simde_uint128, simde__m128i> final
|
||
{
|
||
auto operator()(const simde__m128i & lhs, const simde__m128i & rhs) const
|
||
{
|
||
simde__m128i ret;
|
||
simde_uint128 lhs_, rhs_;
|
||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_uint128));
|
||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_uint128));
|
||
simde_uint128 sum = lhs_ + rhs_;
|
||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m128i));
|
||
return ret;
|
||
}
|
||
};
|
||
|
||
template <> struct add_t<bool, simde__m256i> final : public detail::add32x8_t {};
|
||
// template <> struct add_t<unsigned char, simde__m256i> final : public detail::add32x8_t {};
|
||
template <> struct add_t<psnip_int8_t, simde__m256i> final : public detail::add32x8_t {};
|
||
template <> struct add_t<psnip_uint8_t, simde__m256i> final : public detail::add32x8_t {};
|
||
|
||
template <> struct add_t<psnip_int16_t, simde__m256i> final : public detail::add16x16_t {};
|
||
template <> struct add_t<psnip_uint16_t, simde__m256i> final : public detail::add16x16_t {};
|
||
|
||
template <> struct add_t<psnip_int32_t, simde__m256i> final : public detail::add8x32_t {};
|
||
template <> struct add_t<psnip_uint32_t, simde__m256i> final : public detail::add8x32_t {};
|
||
|
||
template <> struct add_t<psnip_int64_t, simde__m256i> final : public detail::add4x64_t {};
|
||
template <> struct add_t<psnip_uint64_t, simde__m256i> final : public detail::add4x64_t {};
|
||
|
||
template <> struct add_t<simde_int128, simde__m256i> final
|
||
{
|
||
auto operator()(const simde__m256i & lhs, const simde__m256i & rhs) const
|
||
{
|
||
simde__m256i ret;
|
||
simde_int128 lhs_[2], rhs_[2];
|
||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_int128) * 2);
|
||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_int128) * 2);
|
||
simde_int128 sum[2] = { lhs_[0] + rhs_[0], lhs_[1] + rhs_[1] };
|
||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m256i));
|
||
return ret;
|
||
}
|
||
};
|
||
template <> struct add_t<simde_uint128, simde__m256i> final
|
||
{
|
||
auto operator()(const simde__m256i & lhs, const simde__m256i & rhs) const
|
||
{
|
||
simde__m256i ret;
|
||
simde_uint128 lhs_[2], rhs_[2];
|
||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_uint128) * 2);
|
||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_uint128) * 2);
|
||
simde_uint128 sum[2] = { lhs_[0] + rhs_[0], lhs_[1] + rhs_[1] };
|
||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m256i));
|
||
return ret;
|
||
}
|
||
};
|
||
|
||
template <typename OutputT, typename NodeT, std::size_t N> struct add_t<OutputT, std::array<NodeT, N>> final : public detail::add_array_t<OutputT> {};
|
||
template <std::size_t Nbits, typename WordT> struct add_t<dpf::bitstring<Nbits, WordT>, void> final : public detail::bitstring_xor_t<Nbits, WordT> {};
|
||
template <std::size_t N> struct add_t<dpf::blob<N>, void> final : public detail::bitstring_xor_t<N * 8, unsigned char> {};
|
||
template <std::size_t N, typename NodeT>
|
||
struct add_t<dpf::blob<N>, NodeT> final : public detail::bitstring_xor_t<N * 8, unsigned char> {};
|
||
/// @brief Bitwise XOR, not IEEE addition. Float addition does not form an
|
||
/// exact secret-sharing group; XOR of the representation does.
|
||
/// @tparam NodeT GGM node type (not `void`; that selects the scalar wrapper)
|
||
template <typename NodeT>
|
||
struct add_t<float, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||
final : public std::bit_xor<> {};
|
||
template <typename NodeT>
|
||
struct add_t<double, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||
final : public std::bit_xor<> {};
|
||
template <> struct add_t<dpf::bit, void> final : public std::bit_xor<> {};
|
||
template <typename NodeT> struct add_t<dpf::bit, NodeT> final : public std::bit_xor<> {};
|
||
template <typename T> struct add_t<xor_wrapper<T>, void> final : public std::bit_xor<> {};
|
||
// Wildcard unwraps to xor_wrapper<T> + NodeT (SIMD leaf). XOR-group add = bitwise XOR.
|
||
template <typename T, typename NodeT> struct add_t<xor_wrapper<T>, NodeT> final : public std::bit_xor<> {};
|
||
|
||
/// @brief Integer outputs whose width matches a SIMD lane but whose type is
|
||
/// not one of the explicitly specialized aliases (`char`, `long long`,
|
||
/// `char16_t`, and so on).
|
||
/// @tparam OutputT output type
|
||
template <typename OutputT>
|
||
struct add_t<OutputT, simde__m128i,
|
||
std::enable_if_t<std::is_integral_v<OutputT> && sizeof(OutputT) <= 8>>
|
||
: std::conditional_t<sizeof(OutputT) == 1, detail::add16x8_t,
|
||
std::conditional_t<sizeof(OutputT) == 2, detail::add8x16_t,
|
||
std::conditional_t<sizeof(OutputT) == 4, detail::add4x32_t,
|
||
detail::add2x64_t>>> {};
|
||
|
||
template <typename OutputT>
|
||
struct add_t<OutputT, simde__m256i,
|
||
std::enable_if_t<std::is_integral_v<OutputT> && sizeof(OutputT) <= 8>>
|
||
: std::conditional_t<sizeof(OutputT) == 1, detail::add32x8_t,
|
||
std::conditional_t<sizeof(OutputT) == 2, detail::add16x16_t,
|
||
std::conditional_t<sizeof(OutputT) == 4, detail::add8x32_t,
|
||
detail::add4x64_t>>> {};
|
||
|
||
template <std::size_t N>
|
||
struct add_t<dpf::modint<N>, simde__m128i>
|
||
{
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const
|
||
{
|
||
return add_t<typename dpf::modint<N>::integral_type, simde__m128i>{}(a, b);
|
||
}
|
||
};
|
||
|
||
template <std::size_t N>
|
||
struct add_t<dpf::modint<N>, simde__m256i>
|
||
{
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const
|
||
{
|
||
return add_t<typename dpf::modint<N>::integral_type, simde__m256i>{}(a, b);
|
||
}
|
||
};
|
||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||
|
||
namespace detail
|
||
{
|
||
|
||
// /// @brief subtracts vectors of 8-bit integral types
|
||
// template <typename NodeT> struct sub8_t;
|
||
// /// @brief subtracts vectors of 16-bit integral types
|
||
// template <typename NodeT> struct sub16_t;
|
||
// /// @brief subtracts vectors of 32-bit integral types
|
||
// template <typename NodeT> struct sub32_t;
|
||
// /// @brief subtracts vectors of 64-bit integral types
|
||
// template <typename NodeT> struct sub64_t;
|
||
|
||
/// @brief Function object for subtracting vectors of `16x8`-bit integral types
|
||
struct sub16x8_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return simde_mm_sub_epi8(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for subtracting vectors of `8x16`-bit integral types
|
||
struct sub8x16_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return simde_mm_sub_epi16(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for subtracting vectors of `4x32`-bit integral types
|
||
struct sub4x32_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return simde_mm_sub_epi32(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for subtracting vectors of `2x64`-bit integral types
|
||
struct sub2x64_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return simde_mm_sub_epi64(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for subtracting vectors of `32x8`-bit integral types
|
||
struct sub32x8_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return simde_mm256_sub_epi8(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for subtracting vectors of `16x16`-bit integral types
|
||
struct sub16x16_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return simde_mm256_sub_epi16(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for subtracting vectors of `8x32`-bit integral types
|
||
struct sub8x32_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return simde_mm256_sub_epi32(a, b);
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for subtracting vectors of `4x64`-bit integral types
|
||
struct sub4x64_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return simde_mm256_sub_epi64(a, b);
|
||
}
|
||
};
|
||
|
||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||
template <typename OutputT, typename Enabled = void>
|
||
struct sub_array_t
|
||
{
|
||
template <typename T, std::size_t N>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const std::array<T, N> & a, const std::array<T, N> & b) const noexcept
|
||
{
|
||
std::array<T, N> c;
|
||
std::transform(std::begin(a), std::end(a), std::begin(b), std::begin(c),
|
||
[](const T & a, const T & b) { return std::bit_xor<>{}(a, b); });
|
||
return c;
|
||
}
|
||
};
|
||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||
|
||
template <typename OutputT>
|
||
struct sub_array_t<OutputT, std::enable_if_t<dpf::utils::has_operators_plus_minus_v<OutputT>, void>>
|
||
{
|
||
using output_type = OutputT;
|
||
template <typename T, std::size_t N>
|
||
auto operator()(const std::array<T, N> & a, const std::array<T, N> & b) const
|
||
{
|
||
static_assert(sizeof(output_type) == sizeof(std::array<T, N>),
|
||
"arithmetic leaf array and output type must be the same size");
|
||
std::array<T, N> c;
|
||
output_type a_, b_;
|
||
utils::raw_memcpy(&a_, std::data(a), sizeof(a_));
|
||
utils::raw_memcpy(&b_, std::data(b), sizeof(b_));
|
||
output_type c_ = a_ - b_;
|
||
utils::raw_memcpy(std::data(c), &c_, sizeof(c_));
|
||
return c;
|
||
}
|
||
};
|
||
|
||
} // namespace detail
|
||
|
||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||
|
||
template <> struct subtract_t<bool, simde__m128i> final : public detail::sub16x8_t {};
|
||
template <> struct subtract_t<char, simde__m128i> final : public detail::sub16x8_t {};
|
||
// template <> struct subtract_t<unsigned char, simde__m128i> final : public detail::sub16x8_t {};
|
||
template <> struct subtract_t<psnip_int8_t, simde__m128i> final : public detail::sub16x8_t {};
|
||
template <> struct subtract_t<psnip_uint8_t, simde__m128i> final : public detail::sub16x8_t {};
|
||
|
||
template <> struct subtract_t<psnip_int16_t, simde__m128i> final : public detail::sub8x16_t {};
|
||
template <> struct subtract_t<psnip_uint16_t, simde__m128i> final : public detail::sub8x16_t {};
|
||
|
||
template <> struct subtract_t<psnip_int32_t, simde__m128i> final : public detail::sub4x32_t {};
|
||
template <> struct subtract_t<psnip_uint32_t, simde__m128i> final : public detail::sub4x32_t {};
|
||
|
||
template <> struct subtract_t<psnip_int64_t, simde__m128i> final : public detail::sub2x64_t {};
|
||
template <> struct subtract_t<psnip_uint64_t, simde__m128i> final : public detail::sub2x64_t {};
|
||
|
||
template <> struct subtract_t<simde_int128, simde__m128i> final
|
||
{
|
||
auto operator()(const simde__m128i & lhs, const simde__m128i & rhs) const
|
||
{
|
||
simde__m128i ret;
|
||
simde_int128 lhs_, rhs_;
|
||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_int128));
|
||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_int128));
|
||
simde_int128 sum = lhs_ - rhs_;
|
||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m128i));
|
||
return ret;
|
||
}
|
||
};
|
||
template <> struct subtract_t<simde_uint128, simde__m128i> final
|
||
{
|
||
auto operator()(const simde__m128i & lhs, const simde__m128i & rhs) const
|
||
{
|
||
simde__m128i ret;
|
||
simde_uint128 lhs_, rhs_;
|
||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_uint128));
|
||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_uint128));
|
||
simde_uint128 sum = lhs_ - rhs_;
|
||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m128i));
|
||
return ret;
|
||
}
|
||
};
|
||
|
||
template <> struct subtract_t<bool, simde__m256i> final : public detail::sub32x8_t {};
|
||
template <> struct subtract_t<char, simde__m256i> final : public detail::sub32x8_t {};
|
||
// template <> struct subtract_t<unsigned char, simde__m256i> final : public detail::sub32x8_t {};
|
||
template <> struct subtract_t<psnip_int8_t, simde__m256i> final : public detail::sub32x8_t {};
|
||
template <> struct subtract_t<psnip_uint8_t, simde__m256i> final : public detail::sub32x8_t {};
|
||
|
||
template <> struct subtract_t<psnip_int16_t, simde__m256i> final : public detail::sub16x16_t {};
|
||
template <> struct subtract_t<psnip_uint16_t, simde__m256i> final : public detail::sub16x16_t {};
|
||
|
||
template <> struct subtract_t<psnip_int32_t, simde__m256i> final : public detail::sub8x32_t {};
|
||
template <> struct subtract_t<psnip_uint32_t, simde__m256i> final : public detail::sub8x32_t {};
|
||
|
||
template <> struct subtract_t<psnip_int64_t, simde__m256i> final : public detail::sub4x64_t {};
|
||
template <> struct subtract_t<psnip_uint64_t, simde__m256i> final : public detail::sub4x64_t {};
|
||
|
||
template <> struct subtract_t<simde_int128, simde__m256i> final
|
||
{
|
||
auto operator()(const simde__m256i & lhs, const simde__m256i & rhs) const
|
||
{
|
||
simde__m256i ret;
|
||
simde_int128 lhs_[2], rhs_[2];
|
||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_int128) * 2);
|
||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_int128) * 2);
|
||
simde_int128 sum[2] = { lhs_[0] - rhs_[0], lhs_[1] - rhs_[1] };
|
||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m256i));
|
||
return ret;
|
||
}
|
||
};
|
||
template <> struct subtract_t<simde_uint128, simde__m256i> final
|
||
{
|
||
auto operator()(const simde__m256i & lhs, const simde__m256i & rhs) const
|
||
{
|
||
simde__m256i ret;
|
||
simde_uint128 lhs_[2], rhs_[2];
|
||
utils::raw_memcpy(&lhs_, &lhs, sizeof(simde_uint128) * 2);
|
||
utils::raw_memcpy(&rhs_, &rhs, sizeof(simde_uint128) * 2);
|
||
simde_uint128 sum[2] = { lhs_[0] - rhs_[0], lhs_[1] - rhs_[1] };
|
||
utils::raw_memcpy(&ret, &sum, sizeof(simde__m256i));
|
||
return ret;
|
||
}
|
||
};
|
||
|
||
template <typename OutputT, typename NodeT, std::size_t N> struct subtract_t<OutputT, std::array<NodeT, N>> final : public detail::sub_array_t<OutputT> {};
|
||
template <std::size_t Nbits, typename WordT> struct subtract_t<dpf::bitstring<Nbits, WordT>, void> final : public detail::bitstring_xor_t<Nbits, WordT> {};
|
||
template <std::size_t N> struct subtract_t<dpf::blob<N>, void> final : public detail::bitstring_xor_t<N * 8, unsigned char> {};
|
||
template <std::size_t N, typename NodeT>
|
||
struct subtract_t<dpf::blob<N>, NodeT> final : public detail::bitstring_xor_t<N * 8, unsigned char> {};
|
||
/// @brief Bitwise XOR, not IEEE subtraction.
|
||
/// @tparam NodeT GGM node type
|
||
template <typename NodeT>
|
||
struct subtract_t<float, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||
final : public std::bit_xor<> {};
|
||
template <typename NodeT>
|
||
struct subtract_t<double, NodeT, std::enable_if_t<!std::is_void_v<NodeT>>>
|
||
final : public std::bit_xor<> {};
|
||
template <typename NodeT> struct subtract_t<dpf::bit, NodeT> final : public std::bit_xor<> {};
|
||
template <> struct subtract_t<dpf::bit, void> final : public std::bit_xor<> {};
|
||
template <typename T> struct subtract_t<xor_wrapper<T>, void> final : public std::bit_xor<> {};
|
||
// Wildcard unwraps to xor_wrapper<T> + NodeT (SIMD leaf). XOR-group sub = bitwise XOR.
|
||
template <typename T, typename NodeT> struct subtract_t<xor_wrapper<T>, NodeT> final : public std::bit_xor<> {};
|
||
|
||
template <typename OutputT>
|
||
struct subtract_t<OutputT, simde__m128i,
|
||
std::enable_if_t<std::is_integral_v<OutputT> && sizeof(OutputT) <= 8>>
|
||
: std::conditional_t<sizeof(OutputT) == 1, detail::sub16x8_t,
|
||
std::conditional_t<sizeof(OutputT) == 2, detail::sub8x16_t,
|
||
std::conditional_t<sizeof(OutputT) == 4, detail::sub4x32_t,
|
||
detail::sub2x64_t>>> {};
|
||
|
||
template <typename OutputT>
|
||
struct subtract_t<OutputT, simde__m256i,
|
||
std::enable_if_t<std::is_integral_v<OutputT> && sizeof(OutputT) <= 8>>
|
||
: std::conditional_t<sizeof(OutputT) == 1, detail::sub32x8_t,
|
||
std::conditional_t<sizeof(OutputT) == 2, detail::sub16x16_t,
|
||
std::conditional_t<sizeof(OutputT) == 4, detail::sub8x32_t,
|
||
detail::sub4x64_t>>> {};
|
||
|
||
template <std::size_t N>
|
||
struct subtract_t<dpf::modint<N>, simde__m128i>
|
||
{
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const
|
||
{
|
||
return subtract_t<typename dpf::modint<N>::integral_type, simde__m128i>{}(a, b);
|
||
}
|
||
};
|
||
|
||
template <std::size_t N>
|
||
struct subtract_t<dpf::modint<N>, simde__m256i>
|
||
{
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const
|
||
{
|
||
return subtract_t<typename dpf::modint<N>::integral_type, simde__m256i>{}(a, b);
|
||
}
|
||
};
|
||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||
|
||
namespace detail
|
||
{
|
||
|
||
/// @brief Function object for multiplying a vector of `16x8`-bit integral types by an unsigned scalar of the same size
|
||
struct mul16x8_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, psnip_uint8_t b) const noexcept
|
||
{
|
||
auto bb = simde_mm_set1_epi8(b);
|
||
auto lo_bytes = simde_mm_mullo_epi16(a, bb);
|
||
auto hi_bytes = simde_mm_mullo_epi16(simde_mm_srli_epi16(a, 8),
|
||
simde_mm_srli_epi16(bb, 8));
|
||
|
||
return simde_mm_or_si128(
|
||
simde_mm_slli_epi16(hi_bytes, 8),
|
||
simde_mm_and_si128(lo_bytes, simde_mm_set1_epi16(0xff)));
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for multiplying a vector of `8x16`-bit integral types by an unsigned scalar of the same size
|
||
struct mul8x16_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, psnip_uint16_t b) const noexcept
|
||
{
|
||
return simde_mm_mullo_epi16(a, simde_mm_set1_epi16(b));
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for multiplying a vector of `4x32`-bit integral types by an unsigned scalar of the same size
|
||
struct mul4x32_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, psnip_uint32_t b) const noexcept
|
||
{
|
||
return simde_mm_mullo_epi32(a, simde_mm_set1_epi32(b));
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for multiplying a vector of `2x64`-bit integral types by an unsigned scalar of the same size
|
||
struct mul2x64_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, psnip_uint64_t b) const noexcept
|
||
{
|
||
return simde__m128i{static_cast<int64_t>(a[0]*b),
|
||
static_cast<int64_t>(a[1]*b)};
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for multiplying a vector of `32x8`-bit integral types by an unsigned scalar of the same size
|
||
struct mul32x8_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, psnip_uint8_t b) const noexcept
|
||
{
|
||
auto bb = simde_mm256_set1_epi8(b);
|
||
auto lo_bytes = simde_mm256_mullo_epi16(a, bb);
|
||
auto hi_bytes = simde_mm256_mullo_epi16(simde_mm256_srli_epi16(a, 8),
|
||
simde_mm256_srli_epi16(bb, 8));
|
||
|
||
return simde_mm256_or_si256(
|
||
simde_mm256_slli_epi16(hi_bytes, 8),
|
||
simde_mm256_and_si256(lo_bytes, simde_mm256_set1_epi16(0xff)));
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for multiplying a vector of `16x16`-bit integral types by an unsigned scalar of the same size
|
||
struct mul16x16_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, psnip_uint16_t b) const noexcept
|
||
{
|
||
return simde_mm256_mullo_epi16(a, simde_mm256_set1_epi16(b));
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for multiplying a vector of `8x32`-bit integral types by an unsigned scalar of the same size
|
||
struct mul8x32_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, psnip_uint32_t b) const noexcept
|
||
{
|
||
return simde_mm256_mullo_epi32(a, simde_mm256_set1_epi32(b));
|
||
}
|
||
};
|
||
|
||
/// @brief Function object for multiplying a vector of `4x64`-bit integral types by an unsigned scalar of the same size
|
||
struct mul4x64_t
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, psnip_uint64_t b) const noexcept
|
||
{
|
||
return simde__m256i{static_cast<int64_t>(a[0]*b),
|
||
static_cast<int64_t>(a[1]*b),
|
||
static_cast<int64_t>(a[2]*b),
|
||
static_cast<int64_t>(a[3]*b)};
|
||
}
|
||
};
|
||
|
||
} // namespace detail
|
||
|
||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
|
||
|
||
template <> struct multiply_t<bool, simde__m128i> final : public detail::mul16x8_t {};
|
||
template <> struct multiply_t<char, simde__m128i> final : public detail::mul16x8_t {};
|
||
// template <> struct multiply_t<unsigned char, simde__m128i> final : public detail::mul16x8_t {};
|
||
template <> struct multiply_t<psnip_int8_t, simde__m128i> final : public detail::mul16x8_t {};
|
||
template <> struct multiply_t<psnip_uint8_t, simde__m128i> final : public detail::mul16x8_t {};
|
||
|
||
template <> struct multiply_t<psnip_int16_t, simde__m128i> final : public detail::mul8x16_t {};
|
||
template <> struct multiply_t<psnip_uint16_t, simde__m128i> final : public detail::mul8x16_t {};
|
||
|
||
template <> struct multiply_t<psnip_int32_t, simde__m128i> final : public detail::mul4x32_t {};
|
||
template <> struct multiply_t<psnip_uint32_t, simde__m128i> final : public detail::mul4x32_t {};
|
||
|
||
template <> struct multiply_t<psnip_int64_t, simde__m128i> final : public detail::mul2x64_t {};
|
||
template <> struct multiply_t<psnip_uint64_t, simde__m128i> final : public detail::mul2x64_t {};
|
||
|
||
template <> struct multiply_t<bool, simde__m256i> final : public detail::mul32x8_t {};
|
||
// template <> struct multiply_t<unsigned char, simde__m256i> final : public detail::mul32x8_t {};
|
||
template <> struct multiply_t<psnip_int8_t, simde__m256i> final : public detail::mul32x8_t {};
|
||
template <> struct multiply_t<psnip_uint8_t, simde__m256i> final : public detail::mul32x8_t {};
|
||
|
||
template <> struct multiply_t<psnip_int16_t, simde__m256i> final : public detail::mul16x16_t {};
|
||
template <> struct multiply_t<psnip_uint16_t, simde__m256i> final : public detail::mul16x16_t {};
|
||
|
||
template <> struct multiply_t<psnip_int32_t, simde__m256i> final : public detail::mul8x32_t {};
|
||
template <> struct multiply_t<psnip_uint32_t, simde__m256i> final : public detail::mul8x32_t {};
|
||
|
||
template <> struct multiply_t<psnip_int64_t, simde__m256i> final : public detail::mul4x64_t {};
|
||
template <> struct multiply_t<psnip_uint64_t, simde__m256i> final : public detail::mul4x64_t {};
|
||
|
||
template <> struct multiply_t<simde_int128, simde__m128i> final
|
||
{
|
||
auto operator()(const simde__m128i & a, simde_int128 b) const
|
||
{
|
||
simde_int128 a_;
|
||
simde__m128i c;
|
||
utils::raw_memcpy(&a_, &a, sizeof(simde_int128));
|
||
simde_int128 c_ = a_ * b;
|
||
utils::raw_memcpy(&c, &c_, sizeof(simde__m128i));
|
||
return c;
|
||
}
|
||
};
|
||
|
||
template <> struct multiply_t<simde_uint128, simde__m128i> final
|
||
{
|
||
auto operator()(const simde__m128i & a, simde_uint128 b) const
|
||
{
|
||
simde_uint128 a_;
|
||
simde__m128i c;
|
||
utils::raw_memcpy(&a_, &a, sizeof(simde_uint128));
|
||
simde_uint128 c_ = a_ * b;
|
||
utils::raw_memcpy(&c, &c_, sizeof(simde__m128i));
|
||
return c;
|
||
}
|
||
};
|
||
|
||
// template <typename NodeT> struct multiply_t<simde_uint128, NodeT> final : public std::multiplies<simde_uint128> {};
|
||
|
||
// template <typename NodeT> struct multiply_t<float, NodeT> final : public std::bit_and<> {};
|
||
// template <typename NodeT> struct multiply_t<double, NodeT> final : public std::bit_and<> {};
|
||
template <> struct multiply_t<dpf::bit, simde__m128i> final
|
||
{
|
||
auto operator()(const simde__m128i & a, const dpf::bit & b) const
|
||
{
|
||
simde__m128i bb = simde_mm_set1_epi8(-b);
|
||
return simde_mm_and_si128(a, bb);
|
||
}
|
||
};
|
||
template <> struct multiply_t<dpf::bit, simde__m256i> final
|
||
{
|
||
auto operator()(const simde__m256i & a, const dpf::bit & b) const
|
||
{
|
||
simde__m256i bb = simde_mm256_set1_epi8(-b);
|
||
return simde_mm256_and_si256(a, bb);
|
||
}
|
||
};
|
||
// XOR-group scale of a packed leaf: AND each lane with scalar b.
|
||
// std::bit_and<xor_wrapper<T>> is the wrong signature (two wrappers, not NodeT×wrapper).
|
||
template <typename T>
|
||
struct multiply_t<xor_wrapper<T>, simde__m128i> final
|
||
{
|
||
auto operator()(const simde__m128i & a, const xor_wrapper<T> & b) const
|
||
{
|
||
using val_t = typename xor_wrapper<T>::value_type;
|
||
val_t v = static_cast<val_t>(b);
|
||
simde__m128i bb;
|
||
if constexpr (sizeof(val_t) == 8) {
|
||
bb = simde_mm_set1_epi64x(static_cast<int64_t>(v));
|
||
} else if constexpr (sizeof(val_t) == 4) {
|
||
bb = simde_mm_set1_epi32(static_cast<int32_t>(v));
|
||
} else if constexpr (sizeof(val_t) == 2) {
|
||
bb = simde_mm_set1_epi16(static_cast<int16_t>(v));
|
||
} else if constexpr (sizeof(val_t) == 1) {
|
||
bb = simde_mm_set1_epi8(static_cast<int8_t>(v));
|
||
} else {
|
||
alignas(simde__m128i) unsigned char buf[sizeof(simde__m128i)]{};
|
||
for (std::size_t off = 0; off + sizeof(val_t) <= sizeof(buf); off += sizeof(val_t))
|
||
utils::raw_memcpy(buf + off, &v, sizeof(val_t));
|
||
utils::raw_memcpy(&bb, buf, sizeof(bb));
|
||
}
|
||
return simde_mm_and_si128(a, bb);
|
||
}
|
||
};
|
||
template <typename T>
|
||
struct multiply_t<xor_wrapper<T>, simde__m256i> final
|
||
{
|
||
auto operator()(const simde__m256i & a, const xor_wrapper<T> & b) const
|
||
{
|
||
using val_t = typename xor_wrapper<T>::value_type;
|
||
val_t v = static_cast<val_t>(b);
|
||
simde__m256i bb;
|
||
if constexpr (sizeof(val_t) == 8) {
|
||
bb = simde_mm256_set1_epi64x(static_cast<int64_t>(v));
|
||
} else if constexpr (sizeof(val_t) == 4) {
|
||
bb = simde_mm256_set1_epi32(static_cast<int32_t>(v));
|
||
} else if constexpr (sizeof(val_t) == 2) {
|
||
bb = simde_mm256_set1_epi16(static_cast<int16_t>(v));
|
||
} else if constexpr (sizeof(val_t) == 1) {
|
||
bb = simde_mm256_set1_epi8(static_cast<int8_t>(v));
|
||
} else {
|
||
alignas(simde__m256i) unsigned char buf[sizeof(simde__m256i)]{};
|
||
for (std::size_t off = 0; off + sizeof(val_t) <= sizeof(buf); off += sizeof(val_t))
|
||
utils::raw_memcpy(buf + off, &v, sizeof(val_t));
|
||
utils::raw_memcpy(&bb, buf, sizeof(bb));
|
||
}
|
||
return simde_mm256_and_si256(a, bb);
|
||
}
|
||
};
|
||
|
||
template <typename OutputT>
|
||
struct multiply_t<OutputT, simde__m128i,
|
||
std::enable_if_t<std::is_integral_v<OutputT> && sizeof(OutputT) <= 8>>
|
||
: std::conditional_t<sizeof(OutputT) == 1, detail::mul16x8_t,
|
||
std::conditional_t<sizeof(OutputT) == 2, detail::mul8x16_t,
|
||
std::conditional_t<sizeof(OutputT) == 4, detail::mul4x32_t,
|
||
detail::mul2x64_t>>> {};
|
||
|
||
template <typename OutputT>
|
||
struct multiply_t<OutputT, simde__m256i,
|
||
std::enable_if_t<std::is_integral_v<OutputT> && sizeof(OutputT) <= 8>>
|
||
: std::conditional_t<sizeof(OutputT) == 1, detail::mul32x8_t,
|
||
std::conditional_t<sizeof(OutputT) == 2, detail::mul16x16_t,
|
||
std::conditional_t<sizeof(OutputT) == 4, detail::mul8x32_t,
|
||
detail::mul4x64_t>>> {};
|
||
|
||
template <std::size_t N>
|
||
struct multiply_t<dpf::modint<N>, simde__m128i>
|
||
{
|
||
auto operator()(const simde__m128i & a, dpf::modint<N> b) const
|
||
{
|
||
using integral_type = typename dpf::modint<N>::integral_type;
|
||
return multiply_t<integral_type, simde__m128i>{}(a, static_cast<integral_type>(b));
|
||
}
|
||
};
|
||
|
||
template <std::size_t N>
|
||
struct multiply_t<dpf::modint<N>, simde__m256i>
|
||
{
|
||
auto operator()(const simde__m256i & a, dpf::modint<N> b) const
|
||
{
|
||
using integral_type = typename dpf::modint<N>::integral_type;
|
||
return multiply_t<integral_type, simde__m256i>{}(a, static_cast<integral_type>(b));
|
||
}
|
||
};
|
||
|
||
/// @brief `float` and `double` leaves form a bitwise group so shares
|
||
/// reconstruct exactly. Addition is XOR; scaling is AND of the
|
||
/// IEEE bit pattern, not IEEE arithmetic.
|
||
template <>
|
||
struct multiply_t<float, simde__m128i> final
|
||
{
|
||
auto operator()(const simde__m128i & a, float b) const
|
||
{
|
||
static_assert(sizeof(float) == 4, "float must be 32 bits");
|
||
psnip_uint32_t bits = 0;
|
||
utils::raw_memcpy(&bits, &b, sizeof(bits));
|
||
return simde_mm_and_si128(a, simde_mm_set1_epi32(static_cast<int>(bits)));
|
||
}
|
||
};
|
||
|
||
template <>
|
||
struct multiply_t<float, simde__m256i> final
|
||
{
|
||
auto operator()(const simde__m256i & a, float b) const
|
||
{
|
||
static_assert(sizeof(float) == 4, "float must be 32 bits");
|
||
psnip_uint32_t bits = 0;
|
||
utils::raw_memcpy(&bits, &b, sizeof(bits));
|
||
return simde_mm256_and_si256(a, simde_mm256_set1_epi32(static_cast<int>(bits)));
|
||
}
|
||
};
|
||
|
||
template <>
|
||
struct multiply_t<double, simde__m128i> final
|
||
{
|
||
auto operator()(const simde__m128i & a, double b) const
|
||
{
|
||
static_assert(sizeof(double) == 8, "double must be 64 bits");
|
||
psnip_uint64_t bits = 0;
|
||
utils::raw_memcpy(&bits, &b, sizeof(bits));
|
||
return simde_mm_and_si128(a, simde_mm_set1_epi64x(static_cast<long long>(bits)));
|
||
}
|
||
};
|
||
|
||
template <>
|
||
struct multiply_t<double, simde__m256i> final
|
||
{
|
||
auto operator()(const simde__m256i & a, double b) const
|
||
{
|
||
static_assert(sizeof(double) == 8, "double must be 64 bits");
|
||
psnip_uint64_t bits = 0;
|
||
utils::raw_memcpy(&bits, &b, sizeof(bits));
|
||
return simde_mm256_and_si256(a, simde_mm256_set1_epi64x(static_cast<long long>(bits)));
|
||
}
|
||
};
|
||
|
||
namespace packed_spec
|
||
{
|
||
|
||
template <typename T>
|
||
struct is_std_array : std::false_type {};
|
||
|
||
template <typename T, std::size_t N>
|
||
struct is_std_array<std::array<T, N>> : std::true_type {};
|
||
|
||
template <typename T>
|
||
static constexpr bool is_simd_leaf_v
|
||
= std::is_same_v<T, simde__m128i> || std::is_same_v<T, simde__m256i>;
|
||
|
||
template <typename T>
|
||
static constexpr bool byte_fallback_v
|
||
= !std::is_void_v<T> && !is_simd_leaf_v<T> && !is_std_array<T>::value;
|
||
|
||
} // namespace packed_spec
|
||
|
||
template <typename NodeT>
|
||
struct add_t<dpf::twobit, NodeT,
|
||
std::enable_if_t<packed_spec::byte_fallback_v<NodeT>>>
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const NodeT & a, const NodeT & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::add_mod2(a, b);
|
||
}
|
||
};
|
||
|
||
template <typename NodeT>
|
||
struct add_t<dpf::nyble, NodeT,
|
||
std::enable_if_t<packed_spec::byte_fallback_v<NodeT>>>
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const NodeT & a, const NodeT & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::add_mod4(a, b);
|
||
}
|
||
};
|
||
|
||
template <>
|
||
struct add_t<dpf::twobit, simde__m128i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::add_epi2(a, b);
|
||
}
|
||
};
|
||
template <>
|
||
struct add_t<dpf::twobit, simde__m256i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::add_epi2(a, b);
|
||
}
|
||
};
|
||
template <>
|
||
struct add_t<dpf::nyble, simde__m128i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::add_epi4(a, b);
|
||
}
|
||
};
|
||
template <>
|
||
struct add_t<dpf::nyble, simde__m256i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::add_epi4(a, b);
|
||
}
|
||
};
|
||
|
||
template <typename Elem, std::size_t N>
|
||
struct add_t<dpf::twobit, std::array<Elem, N>>
|
||
{
|
||
auto operator()(const std::array<Elem, N> & a, const std::array<Elem, N> & b) const
|
||
{
|
||
std::array<Elem, N> c{};
|
||
for (std::size_t i = 0; i < N; ++i)
|
||
c[i] = add_t<dpf::twobit, Elem>{}(a[i], b[i]);
|
||
return c;
|
||
}
|
||
};
|
||
template <typename Elem, std::size_t N>
|
||
struct add_t<dpf::nyble, std::array<Elem, N>>
|
||
{
|
||
auto operator()(const std::array<Elem, N> & a, const std::array<Elem, N> & b) const
|
||
{
|
||
std::array<Elem, N> c{};
|
||
for (std::size_t i = 0; i < N; ++i)
|
||
c[i] = add_t<dpf::nyble, Elem>{}(a[i], b[i]);
|
||
return c;
|
||
}
|
||
};
|
||
|
||
template <typename NodeT>
|
||
struct subtract_t<dpf::twobit, NodeT,
|
||
std::enable_if_t<packed_spec::byte_fallback_v<NodeT>>>
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const NodeT & a, const NodeT & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::sub_mod2(a, b);
|
||
}
|
||
};
|
||
template <typename NodeT>
|
||
struct subtract_t<dpf::nyble, NodeT,
|
||
std::enable_if_t<packed_spec::byte_fallback_v<NodeT>>>
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const NodeT & a, const NodeT & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::sub_mod4(a, b);
|
||
}
|
||
};
|
||
|
||
template <>
|
||
struct subtract_t<dpf::twobit, simde__m128i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::sub_epi2(a, b);
|
||
}
|
||
};
|
||
template <>
|
||
struct subtract_t<dpf::twobit, simde__m256i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::sub_epi2(a, b);
|
||
}
|
||
};
|
||
template <>
|
||
struct subtract_t<dpf::nyble, simde__m128i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, const simde__m128i & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::sub_epi4(a, b);
|
||
}
|
||
};
|
||
template <>
|
||
struct subtract_t<dpf::nyble, simde__m256i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, const simde__m256i & b) const noexcept
|
||
{
|
||
return dpf::lane_arith::sub_epi4(a, b);
|
||
}
|
||
};
|
||
|
||
template <typename Elem, std::size_t N>
|
||
struct subtract_t<dpf::twobit, std::array<Elem, N>>
|
||
{
|
||
auto operator()(const std::array<Elem, N> & a, const std::array<Elem, N> & b) const
|
||
{
|
||
std::array<Elem, N> c{};
|
||
for (std::size_t i = 0; i < N; ++i)
|
||
c[i] = subtract_t<dpf::twobit, Elem>{}(a[i], b[i]);
|
||
return c;
|
||
}
|
||
};
|
||
template <typename Elem, std::size_t N>
|
||
struct subtract_t<dpf::nyble, std::array<Elem, N>>
|
||
{
|
||
auto operator()(const std::array<Elem, N> & a, const std::array<Elem, N> & b) const
|
||
{
|
||
std::array<Elem, N> c{};
|
||
for (std::size_t i = 0; i < N; ++i)
|
||
c[i] = subtract_t<dpf::nyble, Elem>{}(a[i], b[i]);
|
||
return c;
|
||
}
|
||
};
|
||
|
||
template <typename NodeT>
|
||
struct multiply_t<dpf::twobit, NodeT,
|
||
std::enable_if_t<!std::is_void_v<NodeT> && !packed_spec::is_simd_leaf_v<NodeT>>>
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const NodeT & a, dpf::twobit b) const noexcept
|
||
{
|
||
return dpf::lane_arith::mul_mod2(a, b);
|
||
}
|
||
};
|
||
template <typename NodeT>
|
||
struct multiply_t<dpf::nyble, NodeT,
|
||
std::enable_if_t<!std::is_void_v<NodeT> && !packed_spec::is_simd_leaf_v<NodeT>>>
|
||
{
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_NO_THROW
|
||
HEDLEY_PURE
|
||
auto operator()(const NodeT & a, dpf::nyble b) const noexcept
|
||
{
|
||
return dpf::lane_arith::mul_mod4(a, b);
|
||
}
|
||
};
|
||
|
||
template <>
|
||
struct multiply_t<dpf::twobit, simde__m128i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, dpf::twobit b) const noexcept
|
||
{
|
||
return dpf::lane_arith::mul_epi2(a, b);
|
||
}
|
||
};
|
||
template <>
|
||
struct multiply_t<dpf::twobit, simde__m256i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, dpf::twobit b) const noexcept
|
||
{
|
||
return dpf::lane_arith::mul_epi2(a, b);
|
||
}
|
||
};
|
||
template <>
|
||
struct multiply_t<dpf::nyble, simde__m128i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m128i & a, dpf::nyble b) const noexcept
|
||
{
|
||
return dpf::lane_arith::mul_epi4(a, b);
|
||
}
|
||
};
|
||
template <>
|
||
struct multiply_t<dpf::nyble, simde__m256i> final
|
||
{
|
||
HEDLEY_ALWAYS_INLINE HEDLEY_NO_THROW HEDLEY_PURE
|
||
auto operator()(const simde__m256i & a, dpf::nyble b) const noexcept
|
||
{
|
||
return dpf::lane_arith::mul_epi4(a, b);
|
||
}
|
||
};
|
||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||
|
||
} // namespace leaf_arithmetic
|
||
|
||
} // namespace dpf
|
||
|
||
#endif // LIBDPF_INCLUDE_DPF_LEAF_ARITHMETIC_HPP__
|