libdpf/include/dpf/eval_inner_product.hpp

494 lines
17 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/eval_inner_product.hpp
/// @brief Full / interval DPF evaluation that reduces against a public
/// weight vector instead of materializing the output.
/// @details Same interior + batched exterior AES as `eval_interval`, but
/// each packed leaf is multiply-accumulated into a scalar:
/// additive outputs sum `DPF(x) * w[x]`, XOR outputs xor
/// `DPF(x) & w[x]`. A prepared memoizer skips the interior walk
/// so the tree can be expanded before the weights exist.
#ifndef LIBDPF_INCLUDE_DPF_EVAL_INNER_PRODUCT_HPP__
#define LIBDPF_INCLUDE_DPF_EVAL_INNER_PRODUCT_HPP__
#include <array>
#include <cstddef>
#include <cstring>
#include <limits>
#include <stdexcept>
#include <tuple>
#include <type_traits>
#include <utility>
#include "hedley/hedley.h"
#include <portable-snippets/exact-int/exact-int.h>
#include <simde/simde/x86/avx2.h>
#include "dpf/dpf_key.hpp"
#include "dpf/eval_interval.hpp"
#include "dpf/eval_target.hpp"
#include "dpf/leaf_node.hpp"
#include "dpf/twiddle.hpp"
#include "dpf/utils.hpp"
#include "dpf/xor_wrapper.hpp"
namespace dpf
{
namespace internal
{
template <typename T>
struct is_xor_wrapper : std::false_type {};
template <typename T>
struct is_xor_wrapper<dpf::xor_wrapper<T>> : std::true_type {};
template <typename T>
inline constexpr bool is_xor_wrapper_v = is_xor_wrapper<T>::value;
template <typename W>
HEDLEY_ALWAYS_INLINE
auto weight_as_u64(W && w, std::size_t i)
{
return static_cast<psnip_uint64_t>(w[i]);
}
HEDLEY_ALWAYS_INLINE
simde__m128i load_weight_pair_u64(psnip_uint64_t lo, psnip_uint64_t hi)
{
return simde_mm_set_epi64x(static_cast<int64_t>(hi),
static_cast<int64_t>(lo));
}
HEDLEY_ALWAYS_INLINE
simde__m128i mullo_epi64x2(simde__m128i a, simde__m128i b)
{
#if defined(__AVX512DQ__) && defined(__AVX512VL__)
return _mm_mullo_epi64(a, b);
#else
psnip_uint64_t av[2], bv[2];
std::memcpy(av, &a, sizeof(av));
std::memcpy(bv, &b, sizeof(bv));
av[0] *= bv[0];
av[1] *= bv[1];
simde__m128i r;
std::memcpy(&r, av, sizeof(r));
return r;
#endif
}
template <typename NodeT,
typename OutputsTuple,
std::size_t ...Is>
struct ip_prg_range
{
static constexpr std::size_t pos_min
= const_min_size<block_offset_of_leaf_v<Is, NodeT, OutputsTuple>...>::value;
static constexpr std::size_t pos_end
= const_max_size<(block_offset_of_leaf_v<Is, NodeT, OutputsTuple>
+ block_length_of_leaf_v<std::tuple_element_t<Is, OutputsTuple>, NodeT>)...>::value;
static constexpr std::size_t count = pos_end - pos_min;
};
template <typename OutputT>
struct ip_accum
{
using output_type = OutputT;
static constexpr bool xor_mode = is_xor_wrapper_v<OutputT>;
static constexpr bool simd64 = (sizeof(OutputT) == 8);
simde__m128i vacc = simde_mm_setzero_si128();
output_type scalar{};
template <typename LeafT, typename W>
HEDLEY_ALWAYS_INLINE
void mac(const LeafT & leaf, std::size_t base, std::size_t opl, W && w)
{
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
if constexpr (simd64 && std::is_same_v<LeafT, simde__m128i>)
{
if (HEDLEY_LIKELY(opl == 2))
{
simde__m128i ww = load_weight_pair_u64(
weight_as_u64(w, base),
weight_as_u64(w, base + 1));
if constexpr (xor_mode)
{
vacc = simde_mm_xor_si128(vacc,
simde_mm_and_si128(leaf, ww));
}
else
{
vacc = simde_mm_add_epi64(vacc, mullo_epi64x2(leaf, ww));
}
return;
}
}
HEDLEY_PRAGMA(GCC diagnostic pop)
for (std::size_t p = 0; p < opl; ++p)
{
output_type val;
if constexpr (utils::is_packed_subbyte_v<output_type>)
{
val = extract_leaf<std::remove_cv_t<LeafT>, output_type>(leaf, p);
}
else
{
std::memcpy(&val,
reinterpret_cast<const unsigned char *>(std::addressof(leaf))
+ p * sizeof(output_type),
sizeof(val));
}
const auto wt = weight_as_u64(w, base + p);
if constexpr (xor_mode)
{
scalar = output_type{static_cast<typename output_type::value_type>(
static_cast<psnip_uint64_t>(scalar)
^ (static_cast<psnip_uint64_t>(val) & wt))};
}
else if constexpr (utils::is_packed_subbyte_v<output_type>
&& !std::is_same_v<output_type, dpf::bit>)
{
constexpr unsigned mask
= (1u << utils::packed_lane_bits_v<output_type>) - 1u;
const auto wlane = static_cast<output_type>(
static_cast<unsigned>(wt) & mask);
scalar = scalar + val * wlane;
}
else
{
scalar = static_cast<output_type>(
static_cast<psnip_uint64_t>(scalar)
+ static_cast<psnip_uint64_t>(val) * wt);
}
}
}
HEDLEY_ALWAYS_INLINE
output_type finish() const
{
if constexpr (simd64)
{
psnip_uint64_t lanes[2];
std::memcpy(lanes, &vacc, sizeof(lanes));
if constexpr (xor_mode)
{
return output_type{static_cast<typename output_type::value_type>(
(lanes[0] ^ lanes[1])
^ static_cast<psnip_uint64_t>(scalar))};
}
else
{
return output_type{
lanes[0] + lanes[1]
+ static_cast<psnip_uint64_t>(scalar)};
}
}
return scalar;
}
};
template <std::size_t ...Is,
typename DpfKey,
typename Weights,
typename IntervalMemoizer,
typename IntegralT,
std::size_t ...IIs>
void eval_inner_product_exterior(const DpfKey & dpf, IntegralT from_node,
IntegralT to_node, Weights && weights, IntervalMemoizer && memoizer,
std::index_sequence<IIs...>,
std::tuple<ip_accum<typename DpfKey::concrete_output_type<Is>>...> & accs,
std::size_t start = 0)
{
assert_not_wildcard_output<Is...>(dpf);
if (HEDLEY_UNLIKELY(to_node < from_node && to_node != IntegralT{0}))
throw std::runtime_error("to_node<from_node");
using node_type = typename DpfKey::exterior_node;
using outputs_tuple = typename DpfKey::concrete_outputs_tuple;
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
using range = ip_prg_range<node_type, outputs_tuple, Is...>;
HEDLEY_PRAGMA(GCC diagnostic pop)
constexpr std::size_t opl = DpfKey::outputs_per_leaf;
std::size_t nodes_in_interval = static_cast<std::size_t>(to_node - from_node);
// `to_node == 0` is the saturated exclusive end; the subtraction is the
// leaf count. A real inverted range is rejected above.
auto *nodes = memoizer[DpfKey::depth];
auto cws = std::make_tuple(std::get<Is>(dpf.leaf_nodes).get()...);
HEDLEY_PRAGMA(GCC diagnostic push)
HEDLEY_PRAGMA(GCC diagnostic ignored "-Wignored-attributes")
auto apply_masks = [&](std::size_t k, const node_type & node,
const node_type * HEDLEY_RESTRICT masks)
{
auto apply_output = [&](auto out_index, auto buf_index)
{
constexpr std::size_t out_i = decltype(out_index)::value;
constexpr std::size_t buf_i = decltype(buf_index)::value;
using output_type = typename DpfKey::concrete_output_type<out_i>;
using leaf_type = dpf::leaf_node_t<node_type, output_type>;
constexpr auto pos = block_offset_of_leaf_v<out_i, node_type, outputs_tuple>;
leaf_type mask;
std::memcpy(&mask, masks + (pos - range::pos_min), sizeof(leaf_type));
// Subtractive share: CW_if_t − mask so reconstruct(y0, y1) = y0 − y1 = β.
auto leaf = dpf::subtract_leaf<output_type>(
get_if_lo_bit(std::get<buf_i>(cws), node), mask);
std::get<buf_i>(accs).mac(leaf, k * opl, opl,
utils::get<buf_i>(weights));
};
(apply_output(std::integral_constant<std::size_t, Is>{},
std::integral_constant<std::size_t, IIs>{}), ...);
};
std::size_t j = 0, k = start;
if constexpr (range::count == 2 && range::pos_min == 0)
{
for (; j + 4 <= nodes_in_interval; j += 4, k += 4)
{
alignas(node_type) node_type seeds[4];
alignas(node_type) node_type left[4];
alignas(node_type) node_type right[4];
DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 4; ++t)
{
seeds[t] = utils::to_exterior_node<node_type>(
unset_lo_2bits(nodes[j + t]));
}
DpfKey::exterior_prg::eval01_x4(seeds, left, right);
DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 4; ++t)
{
node_type masks[2] = {left[t], right[t]};
apply_masks(k + t, nodes[j + t], masks);
}
}
}
else if constexpr (range::count == 1)
{
const auto pos = static_cast<psnip_uint32_t>(range::pos_min);
for (; j + 8 <= nodes_in_interval; j += 8, k += 8)
{
alignas(node_type) node_type seeds[8];
alignas(node_type) node_type masks[8];
DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 8; ++t)
{
seeds[t] = utils::to_exterior_node<node_type>(
unset_lo_2bits(nodes[j + t]));
}
DpfKey::exterior_prg::eval_x8(seeds, masks, pos);
DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 8; ++t)
{
apply_masks(k + t, nodes[j + t], &masks[t]);
}
}
for (; j + 4 <= nodes_in_interval; j += 4, k += 4)
{
alignas(node_type) node_type seeds[4];
alignas(node_type) node_type masks[4];
DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 4; ++t)
{
seeds[t] = utils::to_exterior_node<node_type>(
unset_lo_2bits(nodes[j + t]));
}
DpfKey::exterior_prg::eval_x4(seeds, masks, pos);
DPF_UNROLL_LOOP
for (std::size_t t = 0; t < 4; ++t)
{
apply_masks(k + t, nodes[j + t], &masks[t]);
}
}
}
DPF_UNROLL_LOOP
for (; j < nodes_in_interval; ++j, ++k)
{
const auto & node = nodes[j];
auto seed = utils::to_exterior_node<node_type>(unset_lo_2bits(node));
std::array<node_type, range::count> masks;
DpfKey::exterior_prg::eval(seed, masks.data(),
static_cast<psnip_uint32_t>(range::count),
static_cast<psnip_uint32_t>(range::pos_min));
apply_masks(k, node, masks.data());
}
HEDLEY_PRAGMA(GCC diagnostic pop)
}
template <typename DpfKey,
typename InputT,
typename IntervalMemoizer>
void eval_prepare_nodes(const DpfKey & dpf, InputT from, InputT to,
IntervalMemoizer && memoizer)
{
using dpf_type = DpfKey;
using integral_type = typename DpfKey::integral_type;
utils::flip_msb_if_signed_integral(from);
utils::flip_msb_if_signed_integral(to);
integral_type from_node = utils::get_from_node<dpf_type>(from);
integral_type to_node = utils::get_to_node<dpf_type>(to);
constexpr auto to_int = utils::to_integral_type<InputT>{};
const bool wraps = utils::interval_wraps(
static_cast<integral_type>(to_int(from)),
static_cast<integral_type>(to_int(to)),
utils::bitlength_of_v<InputT>);
auto segs = utils::split_leaf_nodes(from_node, to_node, dpf.depth, wraps);
// The memoizer keeps one interval. A wrap is two intervals, and walking
// the first clobbers the second, so only a single segment can be cached.
if (segs.n == 1)
{
eval_interval_interior(dpf, segs.seg[0].from_node, segs.seg[0].to_node,
memoizer);
}
}
template <std::size_t ...Is,
typename DpfKey,
typename InputT,
typename Weights,
typename IntervalMemoizer,
std::size_t ...IIs>
auto eval_inner_product_impl(const DpfKey & dpf, InputT from, InputT to,
Weights && weights, IntervalMemoizer && memoizer,
std::index_sequence<IIs...>)
{
using dpf_type = DpfKey;
using integral_type = typename DpfKey::integral_type;
utils::flip_msb_if_signed_integral(from);
utils::flip_msb_if_signed_integral(to);
integral_type from_node = utils::get_from_node<dpf_type>(from);
integral_type to_node = utils::get_to_node<dpf_type>(to);
constexpr auto to_int = utils::to_integral_type<InputT>{};
const bool wraps = utils::interval_wraps(
static_cast<integral_type>(to_int(from)),
static_cast<integral_type>(to_int(to)),
utils::bitlength_of_v<InputT>);
auto segs = utils::split_leaf_nodes(from_node, to_node, dpf.depth, wraps);
auto accs = std::make_tuple(
ip_accum<typename DpfKey::concrete_output_type<Is>>{}...);
auto idxs = std::index_sequence<IIs...>{};
std::size_t start = 0;
for (std::size_t s = 0; s < segs.n; ++s)
{
const auto & seg = segs.seg[s];
eval_interval_interior(dpf, seg.from_node, seg.to_node, memoizer);
eval_inner_product_exterior<Is...>(dpf, seg.from_node, seg.to_node,
weights, memoizer, idxs, accs, start);
start += seg.count;
}
if constexpr (sizeof...(Is) == 1)
{
return std::get<0>(accs).finish();
}
else
{
return std::make_tuple(std::get<IIs>(accs).finish()...);
}
}
} // namespace internal
/// @brief Expand the interior tree for `[from, to]`. A wrapping interval is left
/// cold: the memoizer holds one half, and walking the first half of the later
/// inner product would clobber a cached second half. Safe to call before the
/// weight vector exists; a subsequent inner-product on the same memoizer
/// skips the interior AES when the interval did not wrap.
/// @tparam DpfKey DPF key type
/// @tparam InputT input domain type
/// @tparam IntervalMemoizer interval memoizer type
/// @param dpf the DPF key
/// @param from the inclusive start of the range
/// @param to the `to`
/// @param memoizer the memoizer built for this key
template <typename DpfKey,
typename InputT,
typename IntervalMemoizer>
HEDLEY_ALWAYS_INLINE
void eval_prepare_interval(const DpfKey & dpf, InputT from, InputT to,
IntervalMemoizer && memoizer)
{
internal::eval_prepare_nodes(dpf, dpf.offset_x(from), dpf.offset_x(to),
memoizer);
}
template <typename DpfKey,
typename IntervalMemoizer>
HEDLEY_ALWAYS_INLINE
void eval_prepare_full(const DpfKey & dpf, IntervalMemoizer && memoizer)
{
using input_type = typename DpfKey::input_type;
eval_prepare_interval(dpf,
std::numeric_limits<input_type>::min(),
std::numeric_limits<input_type>::max(),
memoizer);
}
/// @brief `sum_x DPF_I(x) * w[x]` (additive) or `xor_x DPF_I(x) & w[x]` (XOR).
/// @details `w[j]` is the weight for the `j`-th output in the interval, matching
/// `eval_interval`'s destination layout. Multiple `Is` take a tuple of
/// weight ranges and return a tuple of accumulators; a single `I` takes
/// one range and returns one accumulator.
/// @tparam I output index
/// @tparam Is is
/// @tparam DpfKey DPF key type
/// @tparam InputT input domain type
/// @tparam Weights weights
/// @tparam IntervalMemoizer interval memoizer type
/// @tparam DpfKey DPF key type
/// @param dpf the DPF key
/// @param from the inclusive start of the range
/// @param to the `to`
/// @param weights the weights
/// @param memoizer the memoizer built for this key
/// @return `sum_x DPF_I(x) * w[x]` (additive) or `xor_x DPF_I(x) & w[x]` (XOR)
template <std::size_t I = 0,
std::size_t ...Is,
typename DpfKey,
typename InputT,
typename Weights,
typename IntervalMemoizer,
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
&& !is_multilevel_key_v<DpfKey>, bool> = true>
HEDLEY_ALWAYS_INLINE
auto eval_inner_product(const DpfKey & dpf, InputT from, InputT to,
Weights && weights, IntervalMemoizer && memoizer)
{
assert_not_wildcard_output<I, Is...>(dpf);
return internal::eval_inner_product_impl<I, Is...>(
dpf, dpf.offset_x(from), dpf.offset_x(to),
weights, memoizer, std::make_index_sequence<1 + sizeof...(Is)>{});
}
template <std::size_t I = 0,
std::size_t ...Is,
typename DpfKey,
typename Weights,
typename IntervalMemoizer,
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
&& !is_multilevel_key_v<DpfKey>, bool> = true>
HEDLEY_ALWAYS_INLINE
auto eval_full_inner_product(const DpfKey & dpf, Weights && weights,
IntervalMemoizer && memoizer)
{
using input_type = typename DpfKey::input_type;
return eval_inner_product<I, Is...>(dpf,
std::numeric_limits<input_type>::min(),
std::numeric_limits<input_type>::max(),
weights, memoizer);
}
} // namespace dpf
#endif // LIBDPF_INCLUDE_DPF_EVAL_INNER_PRODUCT_HPP__