/// @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 #include #include #include #include #include #include #include #include "hedley/hedley.h" #include #include #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 struct is_xor_wrapper : std::false_type {}; template struct is_xor_wrapper> : std::true_type {}; template inline constexpr bool is_xor_wrapper_v = is_xor_wrapper::value; template HEDLEY_ALWAYS_INLINE auto weight_as_u64(W && w, std::size_t i) { return static_cast(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(hi), static_cast(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 struct ip_prg_range { static constexpr std::size_t pos_min = const_min_size...>::value; static constexpr std::size_t pos_end = const_max_size<(block_offset_of_leaf_v + block_length_of_leaf_v, NodeT>)...>::value; static constexpr std::size_t count = pos_end - pos_min; }; template struct ip_accum { using output_type = OutputT; static constexpr bool xor_mode = is_xor_wrapper_v; static constexpr bool simd64 = (sizeof(OutputT) == 8); simde__m128i vacc = simde_mm_setzero_si128(); output_type scalar{}; template 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) 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) { val = extract_leaf, output_type>(leaf, p); } else { std::memcpy(&val, reinterpret_cast(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( static_cast(scalar) ^ (static_cast(val) & wt))}; } else if constexpr (utils::is_packed_subbyte_v && !std::is_same_v) { constexpr unsigned mask = (1u << utils::packed_lane_bits_v) - 1u; const auto wlane = static_cast( static_cast(wt) & mask); scalar = scalar + val * wlane; } else { scalar = static_cast( static_cast(scalar) + static_cast(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( (lanes[0] ^ lanes[1]) ^ static_cast(scalar))}; } else { return output_type{ lanes[0] + lanes[1] + static_cast(scalar)}; } } return scalar; } }; template void eval_inner_product_exterior(const DpfKey & dpf, IntegralT from_node, IntegralT to_node, Weights && weights, IntervalMemoizer && memoizer, std::index_sequence, std::tuple>...> & accs, std::size_t start = 0) { assert_not_wildcard_output(dpf); if (HEDLEY_UNLIKELY(to_node < from_node && to_node != IntegralT{0})) throw std::runtime_error("to_node; constexpr std::size_t opl = DpfKey::outputs_per_leaf; std::size_t nodes_in_interval = static_cast(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(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; using leaf_type = dpf::leaf_node_t; constexpr auto pos = block_offset_of_leaf_v; 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( get_if_lo_bit(std::get(cws), node), mask); std::get(accs).mac(leaf, k * opl, opl, utils::get(weights)); }; (apply_output(std::integral_constant{}, std::integral_constant{}), ...); }; 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( 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(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( 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( 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(unset_lo_2bits(node)); std::array masks; DpfKey::exterior_prg::eval(seed, masks.data(), static_cast(range::count), static_cast(range::pos_min)); apply_masks(k, node, masks.data()); } HEDLEY_PRAGMA(GCC diagnostic pop) } template 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(from); integral_type to_node = utils::get_to_node(to); constexpr auto to_int = utils::to_integral_type{}; const bool wraps = utils::interval_wraps( static_cast(to_int(from)), static_cast(to_int(to)), utils::bitlength_of_v); 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 auto eval_inner_product_impl(const DpfKey & dpf, InputT from, InputT to, Weights && weights, IntervalMemoizer && memoizer, std::index_sequence) { 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(from); integral_type to_node = utils::get_to_node(to); constexpr auto to_int = utils::to_integral_type{}; const bool wraps = utils::interval_wraps( static_cast(to_int(from)), static_cast(to_int(to)), utils::bitlength_of_v); auto segs = utils::split_leaf_nodes(from_node, to_node, dpf.depth, wraps); auto accs = std::make_tuple( ip_accum>{}...); auto idxs = std::index_sequence{}; 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(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(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 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 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::min(), std::numeric_limits::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 && !is_multilevel_key_v, bool> = true> HEDLEY_ALWAYS_INLINE auto eval_inner_product(const DpfKey & dpf, InputT from, InputT to, Weights && weights, IntervalMemoizer && memoizer) { assert_not_wildcard_output(dpf); return internal::eval_inner_product_impl( dpf, dpf.offset_x(from), dpf.offset_x(to), weights, memoizer, std::make_index_sequence<1 + sizeof...(Is)>{}); } template && !is_multilevel_key_v, 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(dpf, std::numeric_limits::min(), std::numeric_limits::max(), weights, memoizer); } } // namespace dpf #endif // LIBDPF_INCLUDE_DPF_EVAL_INNER_PRODUCT_HPP__