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>
1329 lines
50 KiB
C++
1329 lines
50 KiB
C++
/// @file dpf/eval_inner_product.hpp
|
||
/// @brief Full / interval / sequence DPF evaluation that reduces against a
|
||
/// public vector instead of materializing the output.
|
||
/// @details Three local forms share the same interval / sequence domains:
|
||
/// - **Batched leaf walk** (no tag): one output, weights in
|
||
/// `eval_interval` layout, exterior AES batched like that walk,
|
||
/// O(1) accumulator. Cost shape matches `eval_interval` on the
|
||
/// same range (Θ(L) nodes) plus a multiply-add per packed slot.
|
||
/// - **`dpf::paired`** (row-wise): one row of the weight vector per
|
||
/// input. A row is a scalar (one output) or a `tuple` / `array`
|
||
/// zipped with several outputs — leaf slots or an ancestor prefix
|
||
/// plus the leaf — read off one path. Products use `operator*`;
|
||
/// they are summed with `operator+`. Same walk cost as the point
|
||
/// list, plus O(1) arithmetic per selected output per input.
|
||
/// - **`dpf::columns`** (transposed / column-wise): one output,
|
||
/// several weight streams, one accumulator per stream, one walk.
|
||
/// Products are *not* summed across streams. A stream is anything
|
||
/// with `w[i]` or `w(i)`. `dpf::project` maps the share before the
|
||
/// multiply; `dpf::also` sees the unmapped share. Cost is one path
|
||
/// walk plus O(stream count) arithmetic per input.
|
||
///
|
||
/// A single-stream `columns` result matches `paired` on that stream;
|
||
/// a single-output batched leaf walk matches `paired` when the
|
||
/// interval is leaf-aligned (covering-leaf weights equal the clipped
|
||
/// domain points). Unaligned intervals still weight every lane of the
|
||
/// covering leaves — same layout as one-key `eval_interval` buffers /
|
||
/// cohort interval inner products, not the clipped iterable. A sized
|
||
/// weight container shorter than that covering span throws
|
||
/// `std::invalid_argument`. `paired` and `columns` throw the same way
|
||
/// when a sized stream is shorter than the point list or the clipped
|
||
/// interval.
|
||
|
||
#ifndef LIBDPF_INCLUDE_DPF_EVAL_INNER_PRODUCT_HPP__
|
||
#define LIBDPF_INCLUDE_DPF_EVAL_INNER_PRODUCT_HPP__
|
||
|
||
#include <algorithm>
|
||
#include <array>
|
||
#include <iterator>
|
||
#include <vector>
|
||
#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_sequence.hpp"
|
||
#include "dpf/eval_target.hpp"
|
||
#include "dpf/incremental.hpp"
|
||
#include "dpf/leaf_node.hpp"
|
||
#include "dpf/path_memoizer.hpp"
|
||
#include "dpf/twiddle.hpp"
|
||
#include "dpf/utils.hpp"
|
||
#include "dpf/xor_wrapper.hpp"
|
||
|
||
namespace dpf
|
||
{
|
||
|
||
namespace detail_ip_check
|
||
{
|
||
|
||
template <typename T, typename = void>
|
||
struct has_container_size : std::false_type
|
||
{
|
||
};
|
||
|
||
template <typename T>
|
||
struct has_container_size<T,
|
||
std::void_t<decltype(std::size(std::declval<const T &>()))>>
|
||
: std::true_type
|
||
{
|
||
};
|
||
|
||
/// @brief Throw when a sized weight container is shorter than the walk.
|
||
/// Unsizable weights (`w(i)` callables, raw pointers) are left to
|
||
/// the caller.
|
||
template <typename W>
|
||
void require_weight_count(const W & w, std::size_t need, const char * what)
|
||
{
|
||
if constexpr (has_container_size<std::decay_t<W>>::value)
|
||
{
|
||
if (static_cast<std::size_t>(std::size(w)) < need)
|
||
throw std::invalid_argument(what);
|
||
}
|
||
}
|
||
|
||
} // namespace detail_ip_check
|
||
|
||
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()...);
|
||
|
||
if constexpr (DpfKey::is_extractable)
|
||
{
|
||
std::size_t j = 0, k = start;
|
||
for (; j < nodes_in_interval; ++j, ++k)
|
||
{
|
||
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;
|
||
auto leaf = dpf.template traverse_exterior<out_i>(nodes[j]);
|
||
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>{}), ...);
|
||
}
|
||
return;
|
||
}
|
||
|
||
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];
|
||
HEDLEY_PRAGMA(GCC diagnostic push)
|
||
HEDLEY_PRAGMA(GCC diagnostic ignored "-Warray-bounds")
|
||
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]));
|
||
}
|
||
HEDLEY_PRAGMA(GCC diagnostic pop)
|
||
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);
|
||
|
||
constexpr std::size_t opl = DpfKey::outputs_per_leaf;
|
||
const std::size_t lanes = segs.total * opl;
|
||
auto check_one = [&](auto which)
|
||
{
|
||
constexpr std::size_t wi = decltype(which)::value;
|
||
detail_ip_check::require_weight_count(
|
||
utils::get<wi>(weights), lanes,
|
||
"inner product weights are shorter than the covering leaves");
|
||
};
|
||
(check_one(std::integral_constant<std::size_t, IIs>{}), ...);
|
||
|
||
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 lane of the *covering* leaves
|
||
/// of `[from, to]`, matching `eval_interval`'s destination buffer (not the
|
||
/// clipped iterable). 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)
|
||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||
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
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
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)>{});
|
||
}
|
||
|
||
/// @brief Inner product with a VDPF path proof over the same interval nodes.
|
||
/// @details Folds once per BFS node (same transcript as `prove_interval`), then
|
||
/// evaluates. Weights are not mixed into the token.
|
||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||
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
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_inner_product(const DpfKey & dpf, InputT from, InputT to,
|
||
Weights && weights, IntervalMemoizer && memoizer, prove_ref pr)
|
||
{
|
||
static_assert(DpfKey::is_verifiable,
|
||
"eval_inner_product(..., prove(π)): key must carry dpf::verifiable");
|
||
detail::vdpf::init_proof(pr.token, dpf);
|
||
prove_fold_interval(dpf, from, to, pr.token);
|
||
detail::vdpf::fold_output_binding(pr.token, dpf);
|
||
return eval_inner_product<I, Is...>(dpf, from, to,
|
||
std::forward<Weights>(weights),
|
||
std::forward<IntervalMemoizer>(memoizer));
|
||
}
|
||
|
||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||
template <std::size_t I = 0,
|
||
std::size_t ...Is,
|
||
typename DpfKey,
|
||
typename InputT,
|
||
typename Weights,
|
||
std::enable_if_t<looks_like_dpf_key_v<DpfKey>
|
||
&& !is_multilevel_key_v<DpfKey>, bool> = true>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_inner_product(const DpfKey & dpf, InputT from, InputT to,
|
||
Weights && weights, prove_ref pr)
|
||
{
|
||
return eval_inner_product<I, Is...>(dpf, from, to,
|
||
std::forward<Weights>(weights),
|
||
dpf::make_basic_interval_memoizer(dpf, from, to), pr);
|
||
}
|
||
|
||
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
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
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);
|
||
}
|
||
|
||
/// @brief Tag: row-wise zip of several DPF outputs with each weight element.
|
||
struct paired_t
|
||
{
|
||
};
|
||
|
||
inline constexpr paired_t paired{};
|
||
|
||
/// @brief Tag: transposed walk — one output, several weight streams kept apart.
|
||
/// @details Unlike `paired`, products are not summed across streams.
|
||
struct columns_t
|
||
{
|
||
};
|
||
|
||
inline constexpr columns_t columns{};
|
||
|
||
/// @brief Map a leaf share before it is multiplied by column weights.
|
||
template <typename F>
|
||
struct project_fn
|
||
{
|
||
F fn;
|
||
};
|
||
|
||
/// @brief Wrap `fn` as the column projector. `fn` is called as `fn(share)`.
|
||
template <typename F>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
project_fn<std::decay_t<F>> project(F && fn)
|
||
{
|
||
return project_fn<std::decay_t<F>>{std::forward<F>(fn)};
|
||
}
|
||
|
||
/// @brief Observe each unmapped share during a `columns` walk.
|
||
template <typename F>
|
||
struct also_fn
|
||
{
|
||
F fn;
|
||
};
|
||
|
||
/// @brief Wrap `fn` as a `columns` side visit.
|
||
/// @details `fn` is called as `fn(i, x, share)`: list index, domain point,
|
||
/// then the share `project` has not seen.
|
||
template <typename F>
|
||
HEDLEY_ALWAYS_INLINE
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
also_fn<std::decay_t<F>> also(F && fn)
|
||
{
|
||
return also_fn<std::decay_t<F>>{std::forward<F>(fn)};
|
||
}
|
||
|
||
/// @brief Column projector that returns the share unchanged.
|
||
struct identity_project
|
||
{
|
||
template <typename T>
|
||
HEDLEY_ALWAYS_INLINE
|
||
constexpr T operator()(T value) const
|
||
{
|
||
return value;
|
||
}
|
||
};
|
||
|
||
/// @brief Column side visit that ignores its arguments.
|
||
struct noop_also
|
||
{
|
||
template <typename... A>
|
||
HEDLEY_ALWAYS_INLINE
|
||
void operator()(A && ...) const noexcept
|
||
{
|
||
}
|
||
};
|
||
|
||
namespace detail_ip
|
||
{
|
||
|
||
template <typename T, typename = void>
|
||
struct has_public_addends : std::false_type
|
||
{
|
||
};
|
||
|
||
template <typename T>
|
||
struct has_public_addends<T, std::void_t<decltype(std::declval<const T &>().public_addends)>>
|
||
: std::true_type
|
||
{
|
||
};
|
||
|
||
template <typename Row>
|
||
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
|
||
{
|
||
};
|
||
|
||
/// @brief Component `K` of a row. A scalar row pairs with output 0 only.
|
||
template <std::size_t K, typename Row>
|
||
HEDLEY_ALWAYS_INLINE
|
||
decltype(auto) component(Row && row)
|
||
{
|
||
using R = std::decay_t<Row>;
|
||
if constexpr (utils::is_tuple_v<R> || is_std_array<R>::value)
|
||
return std::get<K>(std::forward<Row>(row));
|
||
else
|
||
{
|
||
static_assert(K == 0,
|
||
"a scalar weight pairs with one DPF output; use a tuple or "
|
||
"std::array row for several outputs");
|
||
return std::forward<Row>(row);
|
||
}
|
||
}
|
||
|
||
template <std::size_t I, typename Key>
|
||
HEDLEY_ALWAYS_INLINE
|
||
auto share_at(const Key & key, const typename Key::input_type & tx,
|
||
const typename Key::interior_node & node)
|
||
{
|
||
constexpr auto bits = utils::bitlength_of_v<typename Key::input_type>;
|
||
constexpr auto prefix = Key::meta[I].prefix == 0
|
||
? bits : Key::meta[I].prefix;
|
||
auto lane = detail::incr::lane_input(tx, prefix, bits);
|
||
auto leaf = key.template traverse_exterior<I>(node);
|
||
if constexpr (has_public_addends<Key>::value)
|
||
detail::incr::absorb_public_addend_lane<I>(key, leaf, lane);
|
||
using output_type = typename Key::template concrete_output_type<I>;
|
||
return *make_eval_dpf_output<Key, output_type>(leaf, lane);
|
||
}
|
||
|
||
template <std::size_t... Outs, typename Key, typename Tx, typename Path,
|
||
typename Row, typename Acc, std::size_t... Ks>
|
||
HEDLEY_ALWAYS_INLINE
|
||
void mac_row(const Key & key, const Tx & tx, Path & path, Row && row,
|
||
Acc & acc, std::index_sequence<Ks...>)
|
||
{
|
||
// One exterior expansion per output per leaf/ancestor bucket. Sequential
|
||
// and sorted queries reuse the node already on the path.
|
||
((acc = acc + (share_at<Outs>(key, tx, path[Key::meta[Outs].tree_level])
|
||
* component<Ks>(row))), ...);
|
||
}
|
||
|
||
template <std::size_t... Outs, typename Key, typename Point, typename Rows,
|
||
typename Acc>
|
||
void accumulate_points(const Key & key, Point first, Point last, Rows && rows,
|
||
Acc & acc)
|
||
{
|
||
constexpr std::size_t nout = sizeof...(Outs);
|
||
constexpr std::size_t deepest = std::max({std::size_t{0},
|
||
Key::meta[Outs].tree_level...});
|
||
auto path = make_basic_path_memoizer<Key>();
|
||
std::size_t i = 0;
|
||
for (auto it = first; it != last; ++it, ++i)
|
||
{
|
||
auto tx = key.offset_x(*it);
|
||
utils::flip_msb_if_signed_integral(tx);
|
||
detail::ensure_level(key, tx, path, deepest);
|
||
mac_row<Outs...>(key, tx, path, rows[i], acc,
|
||
std::make_index_sequence<nout>{});
|
||
}
|
||
}
|
||
|
||
template <std::size_t I0, std::size_t... Rest>
|
||
struct pack_first
|
||
{
|
||
static constexpr std::size_t value = I0;
|
||
};
|
||
|
||
template <std::size_t... Outs, typename Key, typename Rows>
|
||
auto accum_type_from(const Key &, Rows && rows)
|
||
{
|
||
using row_type = std::decay_t<decltype(rows[std::size_t{0}])>;
|
||
using y0 = decltype(share_at<pack_first<Outs...>::value>(
|
||
std::declval<const Key &>(),
|
||
std::declval<const typename Key::input_type &>(),
|
||
std::declval<const typename Key::interior_node &>()));
|
||
using w0 = std::decay_t<decltype(component<0>(std::declval<row_type &>()))>;
|
||
using acc_type = decltype(std::declval<y0>() * std::declval<w0>());
|
||
return acc_type{};
|
||
}
|
||
|
||
/// @brief `w[i]` when `w` is a range, otherwise `w(i)`.
|
||
template <typename W>
|
||
HEDLEY_ALWAYS_INLINE
|
||
decltype(auto) weight_at(W && w, std::size_t i)
|
||
{
|
||
if constexpr (std::is_invocable_v<W &, std::size_t>)
|
||
return w(i);
|
||
else
|
||
return w[i];
|
||
}
|
||
|
||
template <typename Weights>
|
||
struct is_column_pack : std::false_type
|
||
{
|
||
};
|
||
|
||
template <typename... Ts>
|
||
struct is_column_pack<std::tuple<Ts...>> : std::true_type
|
||
{
|
||
};
|
||
|
||
template <typename T, std::size_t N>
|
||
struct is_column_pack<std::array<T, N>> : std::true_type
|
||
{
|
||
};
|
||
|
||
template <std::size_t Out, std::size_t K, typename Key, typename Weights,
|
||
typename Proj>
|
||
auto column_accum_type()
|
||
{
|
||
using share_type = decltype(share_at<Out>(
|
||
std::declval<const Key &>(),
|
||
std::declval<const typename Key::input_type &>(),
|
||
std::declval<const typename Key::interior_node &>()));
|
||
using mapped = decltype(std::declval<Proj &>()(std::declval<share_type>()));
|
||
using weight = decltype(weight_at(
|
||
std::get<K>(std::declval<Weights &>()), std::size_t{0}));
|
||
using acc_type = decltype(std::declval<mapped>() * std::declval<weight>());
|
||
return acc_type{};
|
||
}
|
||
|
||
template <std::size_t Out, typename Key, typename Point, typename Weights,
|
||
typename Proj, typename Sink, std::size_t... Ks>
|
||
auto accumulate_columns(const Key & key, Point first, Point last,
|
||
Weights && weights, Proj && proj, Sink && sink,
|
||
std::index_sequence<Ks...>)
|
||
{
|
||
static_assert(sizeof...(Ks) > 0, "columns needs at least one weight stream");
|
||
using acc_tuple = std::tuple<decltype(column_accum_type<Out, Ks, Key,
|
||
std::decay_t<Weights>, std::decay_t<Proj>>())...>;
|
||
acc_tuple acc{};
|
||
if constexpr (std::is_base_of_v<std::random_access_iterator_tag,
|
||
typename std::iterator_traits<Point>::iterator_category>)
|
||
{
|
||
const auto n = static_cast<std::size_t>(std::distance(first, last));
|
||
auto guard = [&](auto which)
|
||
{
|
||
constexpr std::size_t k = decltype(which)::value;
|
||
detail_ip_check::require_weight_count(std::get<k>(weights), n,
|
||
"column weights are shorter than the point list");
|
||
};
|
||
(guard(std::integral_constant<std::size_t, Ks>{}), ...);
|
||
}
|
||
constexpr std::size_t deepest = Key::meta[Out].tree_level;
|
||
auto path = make_basic_path_memoizer<Key>();
|
||
std::size_t i = 0;
|
||
for (auto it = first; it != last; ++it, ++i)
|
||
{
|
||
auto tx = key.offset_x(*it);
|
||
utils::flip_msb_if_signed_integral(tx);
|
||
detail::ensure_level(key, tx, path, deepest);
|
||
auto share = share_at<Out>(key, tx, path[Key::meta[Out].tree_level]);
|
||
sink(i, *it, share);
|
||
auto y = proj(share);
|
||
((std::get<Ks>(acc) = std::get<Ks>(acc)
|
||
+ (y * weight_at(std::get<Ks>(weights), i))), ...);
|
||
}
|
||
return acc;
|
||
}
|
||
|
||
} // namespace detail_ip
|
||
|
||
namespace internal_paired
|
||
{
|
||
|
||
template <typename Input, typename Fn>
|
||
void for_inclusive(Input from, Input to, Fn && fn)
|
||
{
|
||
constexpr auto to_int = utils::to_integral_type<Input>{};
|
||
using integral = decltype(to_int(from));
|
||
const bool wraps = utils::interval_wraps(
|
||
static_cast<integral>(to_int(from)),
|
||
static_cast<integral>(to_int(to)),
|
||
utils::bitlength_of_v<Input>);
|
||
auto step = [&](Input x) { fn(x); };
|
||
if (!wraps)
|
||
{
|
||
for (auto x = from;; ++x)
|
||
{
|
||
step(x);
|
||
if (x == to)
|
||
break;
|
||
}
|
||
return;
|
||
}
|
||
const auto hi = std::numeric_limits<Input>::max();
|
||
const auto lo = std::numeric_limits<Input>::min();
|
||
for (auto x = from;; ++x)
|
||
{
|
||
step(x);
|
||
if (x == hi)
|
||
break;
|
||
}
|
||
for (auto x = lo;; ++x)
|
||
{
|
||
step(x);
|
||
if (x == to)
|
||
break;
|
||
}
|
||
}
|
||
|
||
template <std::size_t... Outs, typename Key, typename Rows>
|
||
auto run_points(const Key & key, const std::vector<typename Key::input_type> & xs,
|
||
Rows && rows)
|
||
{
|
||
assert_not_wildcard_output<Outs...>(key);
|
||
if (xs.size() == 0)
|
||
{
|
||
using acc_type = decltype(detail_ip::accum_type_from<Outs...>(key, rows));
|
||
return acc_type{};
|
||
}
|
||
detail_ip_check::require_weight_count(rows, xs.size(),
|
||
"paired weights are shorter than the point list");
|
||
using acc_type = decltype(detail_ip::accum_type_from<Outs...>(key, rows));
|
||
acc_type acc{};
|
||
detail_ip::accumulate_points<Outs...>(key, xs.begin(), xs.end(), rows, acc);
|
||
return acc;
|
||
}
|
||
|
||
/// @brief Clipped interval dot via the interval tree, not one path per point.
|
||
/// Wrapping intervals stay on the path walk.
|
||
template <std::size_t I, typename DpfKey, typename InputT, typename Rows>
|
||
auto interval_scalar(const DpfKey & dpf, InputT from, InputT to, Rows && rows)
|
||
{
|
||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||
using integral = decltype(to_int(from));
|
||
const bool wraps = utils::interval_wraps(
|
||
static_cast<integral>(to_int(from)),
|
||
static_cast<integral>(to_int(to)),
|
||
utils::bitlength_of_v<InputT>);
|
||
if (wraps)
|
||
{
|
||
std::vector<InputT> xs;
|
||
for_inclusive(from, to, [&](InputT x) { xs.push_back(x); });
|
||
return run_points<I>(dpf, xs, std::forward<Rows>(rows));
|
||
}
|
||
const auto npoints = static_cast<std::size_t>(
|
||
static_cast<integral>(to_int(to)) - static_cast<integral>(to_int(from)))
|
||
+ std::size_t{1};
|
||
detail_ip_check::require_weight_count(rows, npoints,
|
||
"paired weights are shorter than the interval");
|
||
auto buf = make_output_buffer_for_interval<I>(dpf, from, to);
|
||
auto iter = eval_interval<I>(dpf, from, to, buf);
|
||
using share_t = std::decay_t<decltype(*iter.begin())>;
|
||
using weight_t = std::decay_t<decltype(rows[std::size_t{0}])>;
|
||
using acc_t = decltype(std::declval<share_t>() * std::declval<weight_t>());
|
||
acc_t acc{};
|
||
std::size_t i = 0;
|
||
for (auto it = iter.begin(); it != iter.end(); ++it, ++i)
|
||
acc = acc + (*it * rows[i]);
|
||
return acc;
|
||
}
|
||
|
||
} // namespace internal_paired
|
||
|
||
/// @brief `sum_i Σ_k DPF_{out_k}(x_i) * row_i[k]` over `[from, to]`.
|
||
/// @details One row of `rows` per input, in interval order (wrapping the same
|
||
/// way as `eval_interval`). A row is a scalar when one output is
|
||
/// selected, or a `std::tuple` / `std::array` with one component per
|
||
/// output. Ancestor slots and leaf slots are read off the same path.
|
||
/// Differs from the batched leaf walk: one path step per domain point
|
||
/// (reuse via the path memoizer), not batched exterior AES over leaf
|
||
/// nodes. For a single output the opened result matches the batched
|
||
/// form when `rows` is the interval weight vector.
|
||
/// \complexity One path ensure per input up to the deepest selected output,
|
||
/// plus one exterior expand and multiply-add per selected output per
|
||
/// input. Accumulator is O(1); no output buffer.
|
||
template <std::size_t I = 0, std::size_t... Is, typename DpfKey, typename InputT,
|
||
typename Rows>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_inner_product(paired_t, const DpfKey & dpf, InputT from, InputT to,
|
||
Rows && rows)
|
||
{
|
||
using row_t = std::decay_t<decltype(std::declval<Rows &>()[std::size_t{0}])>;
|
||
constexpr bool scalar = !utils::is_tuple_v<row_t>
|
||
&& !detail_ip::is_std_array<row_t>::value;
|
||
if constexpr (scalar && sizeof...(Is) == 0 && !is_multilevel_key_v<DpfKey>)
|
||
{
|
||
return internal_paired::interval_scalar<I>(dpf, from, to,
|
||
std::forward<Rows>(rows));
|
||
}
|
||
else
|
||
{
|
||
std::vector<InputT> xs;
|
||
internal_paired::for_inclusive(from, to, [&](InputT x) { xs.push_back(x); });
|
||
return internal_paired::run_points<I, Is...>(dpf, xs,
|
||
std::forward<Rows>(rows));
|
||
}
|
||
}
|
||
|
||
/// @brief Paired inner product over the whole input domain.
|
||
template <std::size_t I = 0, std::size_t... Is, typename DpfKey, typename Rows>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_full_inner_product(paired_t, const DpfKey & dpf, Rows && rows)
|
||
{
|
||
using input_type = typename DpfKey::input_type;
|
||
return eval_inner_product<I, Is...>(paired, dpf,
|
||
std::numeric_limits<input_type>::min(),
|
||
std::numeric_limits<input_type>::max(),
|
||
std::forward<Rows>(rows));
|
||
}
|
||
|
||
/// @brief Paired inner product over a sorted point list.
|
||
/// @details Same order and sortedness rule as `eval_sequence`. Each point
|
||
/// pairs with `rows[i]`.
|
||
template <std::size_t I = 0, std::size_t... Is, typename DpfKey,
|
||
typename ForwardIterator, typename Rows>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_sequence_inner_product(const DpfKey & dpf, ForwardIterator begin,
|
||
ForwardIterator end, Rows && rows)
|
||
{
|
||
if (!std::is_sorted(begin, end))
|
||
throw std::runtime_error("list must be sorted");
|
||
std::vector<typename DpfKey::input_type> xs(begin, end);
|
||
return internal_paired::run_points<I, Is...>(dpf, xs,
|
||
std::forward<Rows>(rows));
|
||
}
|
||
|
||
/// @brief Paired inner product over a `sequence_recipe`'s points.
|
||
/// @details `points` is the same sorted list the recipe was built from. The
|
||
/// recipe drives nothing the path walk does not already share; it
|
||
/// checks that the list still matches the recipe's output count.
|
||
template <std::size_t I = 0, std::size_t... Is, typename DpfKey,
|
||
typename ForwardIterator, typename Rows>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_sequence_inner_product(const DpfKey & dpf,
|
||
const sequence_recipe & recipe, ForwardIterator begin, ForwardIterator end,
|
||
Rows && rows)
|
||
{
|
||
const auto n = static_cast<std::size_t>(std::distance(begin, end));
|
||
if (n != recipe.output_indices().size())
|
||
throw std::invalid_argument(
|
||
"eval_sequence_inner_product: recipe and point list differ");
|
||
return eval_sequence_inner_product<I, Is...>(dpf, begin, end,
|
||
std::forward<Rows>(rows));
|
||
}
|
||
|
||
namespace detail_columns
|
||
{
|
||
|
||
template <std::size_t I, typename DpfKey, typename ForwardIterator,
|
||
typename Weights, typename Proj, typename Sink>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto run(const DpfKey & dpf, ForwardIterator begin, ForwardIterator end,
|
||
Weights && weights, Proj && proj, Sink && sink, bool require_sorted)
|
||
{
|
||
using pack = std::decay_t<Weights>;
|
||
static_assert(detail_ip::is_column_pack<pack>::value,
|
||
"columns weights are a std::tuple or std::array of streams; "
|
||
"each stream is w[i] or w(i)");
|
||
if (require_sorted && !std::is_sorted(begin, end))
|
||
throw std::runtime_error("list must be sorted");
|
||
std::vector<typename DpfKey::input_type> xs(begin, end);
|
||
constexpr auto n = std::tuple_size<pack>::value;
|
||
return detail_ip::accumulate_columns<I>(dpf, xs.begin(), xs.end(),
|
||
std::forward<Weights>(weights), std::forward<Proj>(proj),
|
||
std::forward<Sink>(sink), std::make_index_sequence<n>{});
|
||
}
|
||
|
||
template <std::size_t I, typename DpfKey, typename InputT, typename Weights,
|
||
std::size_t... Ks>
|
||
auto interval_dots(const DpfKey & dpf, InputT from, InputT to, Weights && weights,
|
||
std::index_sequence<Ks...>)
|
||
{
|
||
constexpr auto to_int = utils::to_integral_type<InputT>{};
|
||
using integral = decltype(to_int(from));
|
||
const bool wraps = utils::interval_wraps(
|
||
static_cast<integral>(to_int(from)),
|
||
static_cast<integral>(to_int(to)),
|
||
utils::bitlength_of_v<InputT>);
|
||
if (wraps)
|
||
{
|
||
std::vector<InputT> xs;
|
||
internal_paired::for_inclusive(from, to, [&](InputT x) { xs.push_back(x); });
|
||
return detail_ip::accumulate_columns<I>(dpf, xs.begin(), xs.end(),
|
||
std::forward<Weights>(weights), identity_project{},
|
||
noop_also{}, std::index_sequence<Ks...>{});
|
||
}
|
||
const auto npoints = static_cast<std::size_t>(
|
||
static_cast<integral>(to_int(to)) - static_cast<integral>(to_int(from)))
|
||
+ std::size_t{1};
|
||
auto guard = [&](auto which)
|
||
{
|
||
constexpr std::size_t k = decltype(which)::value;
|
||
detail_ip_check::require_weight_count(std::get<k>(weights), npoints,
|
||
"column weights are shorter than the interval");
|
||
};
|
||
(guard(std::integral_constant<std::size_t, Ks>{}), ...);
|
||
auto buf = make_output_buffer_for_interval<I>(dpf, from, to);
|
||
auto iter = eval_interval<I>(dpf, from, to, buf);
|
||
using acc_tuple = std::tuple<decltype(detail_ip::column_accum_type<I, Ks,
|
||
DpfKey, std::decay_t<Weights>, identity_project>())...>;
|
||
acc_tuple acc{};
|
||
std::size_t i = 0;
|
||
for (auto it = iter.begin(); it != iter.end(); ++it, ++i)
|
||
{
|
||
auto y = *it;
|
||
((std::get<Ks>(acc) = std::get<Ks>(acc)
|
||
+ (y * detail_ip::weight_at(std::get<Ks>(weights), i))), ...);
|
||
}
|
||
return acc;
|
||
}
|
||
|
||
} // namespace detail_columns
|
||
|
||
/// @brief Several independent dots of one output, one walk (transposed form).
|
||
/// @details `weights` is a `std::tuple` or `std::array` of streams. Stream
|
||
/// `k` is either `w[i]` or `w(i)`, `i` the position in the point
|
||
/// list (or in the interval, wrapping the same way as
|
||
/// `eval_interval`). The result is a tuple of accumulators,
|
||
/// `acc_k = sum_i project(DPF(x_i)) * stream_k(i)`.
|
||
/// `dpf::project(fn)` maps the share first. `dpf::also(fn)` is
|
||
/// called as `fn(i, x, share)` on the unmapped share. Pass either
|
||
/// tag, both, or neither. `also` then `project` is accepted too.
|
||
/// Sequence points are sorted nondecreasing, same as `eval_sequence`.
|
||
/// Relative to `paired`: same path walk for one output, but streams
|
||
/// stay separate (no cross-stream sum). One stream equals `paired`
|
||
/// on that stream. Relative to the batched leaf walk: path-per-point
|
||
/// instead of batched exterior AES; same opened scalar when the
|
||
/// single stream matches the interval weight layout.
|
||
/// \complexity One path ensure and one exterior expand per input, plus
|
||
/// O(stream count) multiply-adds per input. Accumulators are O(stream count).
|
||
/// @{
|
||
|
||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||
typename Weights>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||
ForwardIterator begin, ForwardIterator end, Weights && weights)
|
||
{
|
||
assert_not_wildcard_output<I>(dpf);
|
||
return detail_columns::run<I>(dpf, begin, end,
|
||
std::forward<Weights>(weights), identity_project{}, noop_also{}, true);
|
||
}
|
||
|
||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||
typename Weights, typename Proj>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||
ForwardIterator begin, ForwardIterator end, Weights && weights,
|
||
project_fn<Proj> proj)
|
||
{
|
||
assert_not_wildcard_output<I>(dpf);
|
||
return detail_columns::run<I>(dpf, begin, end,
|
||
std::forward<Weights>(weights), proj.fn, noop_also{}, true);
|
||
}
|
||
|
||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||
typename Weights, typename Sink>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||
ForwardIterator begin, ForwardIterator end, Weights && weights,
|
||
also_fn<Sink> sink)
|
||
{
|
||
assert_not_wildcard_output<I>(dpf);
|
||
return detail_columns::run<I>(dpf, begin, end,
|
||
std::forward<Weights>(weights), identity_project{}, sink.fn, true);
|
||
}
|
||
|
||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||
typename Weights, typename Proj, typename Sink>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||
ForwardIterator begin, ForwardIterator end, Weights && weights,
|
||
project_fn<Proj> proj, also_fn<Sink> sink)
|
||
{
|
||
assert_not_wildcard_output<I>(dpf);
|
||
return detail_columns::run<I>(dpf, begin, end,
|
||
std::forward<Weights>(weights), proj.fn, sink.fn, true);
|
||
}
|
||
|
||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||
typename Weights, typename Proj, typename Sink>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||
ForwardIterator begin, ForwardIterator end, Weights && weights,
|
||
also_fn<Sink> sink, project_fn<Proj> proj)
|
||
{
|
||
return eval_sequence_inner_product<I>(columns, dpf, begin, end,
|
||
std::forward<Weights>(weights), std::move(proj), std::move(sink));
|
||
}
|
||
|
||
template <std::size_t I = 0, typename DpfKey, typename ForwardIterator,
|
||
typename Weights, typename... Extra>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_sequence_inner_product(columns_t, const DpfKey & dpf,
|
||
const sequence_recipe & recipe, ForwardIterator begin,
|
||
ForwardIterator end, Weights && weights, Extra && ... extra)
|
||
{
|
||
const auto n = static_cast<std::size_t>(std::distance(begin, end));
|
||
if (n != recipe.output_indices().size())
|
||
throw std::invalid_argument(
|
||
"eval_sequence_inner_product: recipe and point list differ");
|
||
return eval_sequence_inner_product<I>(columns, dpf, begin, end,
|
||
std::forward<Weights>(weights), std::forward<Extra>(extra)...);
|
||
}
|
||
|
||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights,
|
||
typename Proj, typename Sink>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||
Weights && weights, Proj && proj, Sink && sink)
|
||
{
|
||
assert_not_wildcard_output<I>(dpf);
|
||
if constexpr (std::is_same_v<std::decay_t<Proj>, identity_project>
|
||
&& std::is_same_v<std::decay_t<Sink>, noop_also>
|
||
&& !is_multilevel_key_v<DpfKey>)
|
||
{
|
||
using pack = std::decay_t<Weights>;
|
||
static_assert(detail_ip::is_column_pack<pack>::value,
|
||
"columns weights are a std::tuple or std::array of streams; "
|
||
"each stream is w[i] or w(i)");
|
||
constexpr auto nstreams = std::tuple_size<pack>::value;
|
||
return detail_columns::interval_dots<I>(dpf, from, to,
|
||
std::forward<Weights>(weights), std::make_index_sequence<nstreams>{});
|
||
}
|
||
else
|
||
{
|
||
std::vector<InputT> xs;
|
||
internal_paired::for_inclusive(from, to, [&](InputT x) { xs.push_back(x); });
|
||
return detail_columns::run<I>(dpf, xs.begin(), xs.end(),
|
||
std::forward<Weights>(weights), std::forward<Proj>(proj),
|
||
std::forward<Sink>(sink), false);
|
||
}
|
||
}
|
||
|
||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||
Weights && weights)
|
||
{
|
||
return eval_inner_product<I>(columns, dpf, from, to,
|
||
std::forward<Weights>(weights), identity_project{}, noop_also{});
|
||
}
|
||
|
||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights,
|
||
typename Proj>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||
Weights && weights, project_fn<Proj> proj)
|
||
{
|
||
return eval_inner_product<I>(columns, dpf, from, to,
|
||
std::forward<Weights>(weights), proj.fn, noop_also{});
|
||
}
|
||
|
||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights,
|
||
typename Sink>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||
Weights && weights, also_fn<Sink> sink)
|
||
{
|
||
return eval_inner_product<I>(columns, dpf, from, to,
|
||
std::forward<Weights>(weights), identity_project{}, sink.fn);
|
||
}
|
||
|
||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights,
|
||
typename Proj, typename Sink>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||
Weights && weights, project_fn<Proj> proj, also_fn<Sink> sink)
|
||
{
|
||
return eval_inner_product<I>(columns, dpf, from, to,
|
||
std::forward<Weights>(weights), proj.fn, sink.fn);
|
||
}
|
||
|
||
/// \complexity Same interior expansion as `eval_interval` on `[from, to]` (Θ(L) nodes, L = leaf nodes in the interval) plus a multiply-add per output slot into an O(1) accumulator. The basic memoizer still holds O(L) nodes.
|
||
template <std::size_t I = 0, typename DpfKey, typename InputT, typename Weights,
|
||
typename Proj, typename Sink>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_inner_product(columns_t, const DpfKey & dpf, InputT from, InputT to,
|
||
Weights && weights, also_fn<Sink> sink, project_fn<Proj> proj)
|
||
{
|
||
return eval_inner_product<I>(columns, dpf, from, to,
|
||
std::forward<Weights>(weights), proj.fn, sink.fn);
|
||
}
|
||
|
||
/// @brief `columns` over the whole input domain. Same streams as the interval form.
|
||
template <std::size_t I = 0, typename DpfKey, typename Weights, typename... Extra>
|
||
HEDLEY_WARN_UNUSED_RESULT
|
||
auto eval_full_inner_product(columns_t, const DpfKey & dpf, Weights && weights,
|
||
Extra && ... extra)
|
||
{
|
||
using input_type = typename DpfKey::input_type;
|
||
return eval_inner_product<I>(columns, dpf,
|
||
std::numeric_limits<input_type>::min(),
|
||
std::numeric_limits<input_type>::max(),
|
||
std::forward<Weights>(weights), std::forward<Extra>(extra)...);
|
||
}
|
||
|
||
/// @}
|
||
|
||
} // namespace dpf
|
||
|
||
#endif // LIBDPF_INCLUDE_DPF_EVAL_INNER_PRODUCT_HPP__
|