2026-09-24 14:08:32 -06:00
|
|
|
|
/// @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>)
|
|
|
|
|
|
HEDLEY_PRAGMA(GCC diagnostic pop)
|
|
|
|
|
|
{
|
|
|
|
|
|
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;
|
|
|
|
|
|
}
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
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;
|
|
|
|
|
|
using range = ip_prg_range<node_type, outputs_tuple, Is...>;
|
|
|
|
|
|
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);
|
2026-09-24 20:44:07 -06:00
|
|
|
|
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);
|
2026-09-24 14:08:32 -06:00
|
|
|
|
// 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);
|
2026-09-24 20:44:07 -06:00
|
|
|
|
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);
|
2026-09-24 14:08:32 -06:00
|
|
|
|
|
|
|
|
|
|
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
|
|
|
|
|
|
|
|
|
|
|
|
/// 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.
|
|
|
|
|
|
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);
|
|
|
|
|
|
}
|
|
|
|
|
|
|
|
|
|
|
|
/// `sum_x DPF_I(x) * w[x]` (additive) or `xor_x DPF_I(x) & w[x]` (XOR).
|
|
|
|
|
|
/// `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.
|
|
|
|
|
|
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__
|