libdpf/include/dpf/leaf_arithmetic.hpp
Ryan Henry 0d22946a0e Checkpoint the party/runtime stack before share-program and malicious-mode work.
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>
2026-09-28 05:59:19 -06:00

1379 lines
47 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

/// @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__