m1une's library

This documentation is automatically generated by online-judge-tools/verification-helper

View on GitHub

:heavy_check_mark: Binary Trie with Monoid
(ds/binary_trie/binary_trie_monoid.hpp)

Overview

BinaryTrieMonoid stores pairs (key, value). For a query (x, upper), it returns the monoid product of every stored value whose key satisfies (key ^ x) < upper.

This directly supports queries of the form

prod(B_i) for every i such that (A_i ^ x) < a

by inserting each pair with insert(A_i, B_i) and calling prod_xor_less(x, a).

The monoid operation must be commutative because the selected entries do not have an intrinsic order. Addition, multiplication, minimum, maximum, gcd, and bitwise xor are suitable examples.

Template Parameters

An exclusive upper bound may be larger than the largest representable BitWidth-bit key when that value fits in UInt. For example, with BitWidth == 30, passing 1U << 30 includes every possible key.

Methods

Let $B$ be BitWidth.

Method Description Complexity
BinaryTrieMonoid() Constructs an empty trie. $O(1)$
BinaryTrieMonoid(init) Constructs from (key, value) pairs. $O(NB)$
BinaryTrieMonoid(first, last) Constructs from a range of (key, value) pairs. $O(NB)$
BinaryTrieMonoid(keys, values) Constructs from parallel key and value vectors. $O(NB)$
int size() const Returns the number of inserted pairs. $O(1)$
bool empty() const Returns whether the trie is empty. $O(1)$
node_id root() const Returns the root node handle. $O(1)$
node_id find(UInt key) const Returns the leaf handle for key, or null_node if absent. $O(B)$
const Node& node(node_id id) const Returns a read-only view of a node. $O(1)$
size_t node_count() const Returns allocated nodes, including the root. $O(1)$
void reserve(size_t n) Reserves storage for approximately n nodes. $O(K)$
UInt xor_mask() const Returns the current lazy xor mask. $O(1)$
void clear() Removes every pair and resets the lazy xor. $O(1)$
node_id insert(UInt key, const T& value) Inserts one pair and returns its leaf handle. Duplicate keys are allowed. $O(B)$
int count(UInt key) const Returns the number of pairs with this key. $O(B)$
bool contains(UInt key) const Returns whether this key exists. $O(B)$
T prod(UInt key) const Returns the product of values with exactly this key. $O(B)$
T all_prod() const Returns the product of all stored values. $O(1)$
int erase_all(UInt key) Removes all pairs with this key and returns their count. $O(B)$
void xor_all(UInt value) Applies xor with value to every stored key. $O(1)$

node_id is an integer handle and null_node is its invalid value. A Node exposes child[2], count, and prod. Handles remain valid across insertions and erasures and can index user-owned metadata; clear() invalidates every old handle except the root. References returned by node() may be invalidated by insertion, reserve(), or clear().

The child links describe the physically stored bit paths. After xor_all, logical keys differ by xor_mask(); find(key) accounts for this mask automatically. Here $K$ is the allocated node count. Erasing does not reclaim nodes.

XOR order statistics

Method Description Complexity
UInt kth_xor(int k, UInt x) const Returns the 0-indexed k-th smallest result among key ^ x, including duplicate keys. $O(B)$
UInt kth(int k) const Returns the 0-indexed k-th smallest key. $O(B)$
UInt min() const, UInt max() const Returns the smallest or largest key. Requires a nonempty trie. $O(B)$
UInt min_xor(UInt x) const, UInt max_xor(UInt x) const Returns the minimum or maximum result among key ^ x. Requires a nonempty trie. $O(B)$

XOR counts

Method Description Complexity
int count_xor_equal(UInt x, UInt target) const Counts pairs satisfying (key ^ x) == target. $O(B)$
int count_xor_less(UInt x, UInt upper) const Counts pairs satisfying (key ^ x) < upper. $O(B)$
int count_xor_less_equal(UInt x, UInt upper) const Counts pairs satisfying (key ^ x) <= upper. $O(B)$
int count_xor_greater(UInt x, UInt lower) const Counts pairs satisfying (key ^ x) > lower. $O(B)$
int count_xor_greater_equal(UInt x, UInt lower) const Counts pairs satisfying (key ^ x) >= lower. $O(B)$
int count_xor_range(UInt x, UInt lower, UInt upper) const Counts pairs satisfying lower <= (key ^ x) < upper. $O(B)$

count_less_xor(x, upper) is a compatibility alias for count_xor_less(x, upper). The methods order_of_key, count_less, count_less_equal, count_greater, count_greater_equal, and count_range provide the same comparisons without an xor operand.

XOR products

Method Description Complexity
T prod_xor_equal(UInt x, UInt target) const Returns the product for pairs satisfying (key ^ x) == target. $O(B)$
T prod_xor_less(UInt x, UInt upper) const Returns the product for pairs satisfying (key ^ x) < upper. $O(B)$
T prod_xor_less_equal(UInt x, UInt upper) const Returns the product for pairs satisfying (key ^ x) <= upper. $O(B)$
T prod_xor_greater(UInt x, UInt lower) const Returns the product for pairs satisfying (key ^ x) > lower. $O(B)$
T prod_xor_greater_equal(UInt x, UInt lower) const Returns the product for pairs satisfying (key ^ x) >= lower. $O(B)$
T prod_xor_range(UInt x, UInt lower, UInt upper) const Returns the product for pairs satisfying lower <= (key ^ x) < upper. $O(B)$

When no pair satisfies a product query, it returns Monoid::id(). The methods prod_less, prod_less_equal, prod_greater, prod_greater_equal, and prod_range provide the same comparisons without an xor operand.

Example

#include "ds/binary_trie/binary_trie_monoid.hpp"
#include "monoid/mul.hpp"

#include <cstdint>
#include <iostream>
#include <vector>

int main() {
    std::vector<std::uint32_t> A = {1, 2, 7, 7};
    std::vector<long long> B = {2, 3, 5, 11};

    using Product = m1une::monoid::Mul<long long>;
    m1une::ds::BinaryTrieMonoid<Product, std::uint32_t, 30> trie(A, B);

    std::uint32_t x = 3;
    std::uint32_t a = 4;

    // 1 ^ 3 = 2 and 2 ^ 3 = 1 are less than 4.
    // The answer is B[0] * B[1] = 2 * 3 = 6.
    std::cout << trie.prod_xor_less(x, a) << "\n";

    // Product for 1 <= (A[i] ^ x) < 5.
    std::cout << trie.prod_xor_range(x, 1, 5) << "\n";
}

Depends on

Verified with

Code

#ifndef M1UNE_DS_BINARY_TRIE_MONOID_HPP
#define M1UNE_DS_BINARY_TRIE_MONOID_HPP 1

#include <cassert>
#include <cstdint>
#include <initializer_list>
#include <limits>
#include <type_traits>
#include <utility>
#include <vector>

#include "../../monoid/concept.hpp"

namespace m1une {
namespace ds {

template <m1une::monoid::IsMonoid Monoid,
          typename UInt = std::uint32_t,
          int BitWidth = std::numeric_limits<UInt>::digits>
struct BinaryTrieMonoid {
    using T = typename Monoid::value_type;

    static_assert(std::is_integral_v<UInt>);
    static_assert(std::is_unsigned_v<UInt>);
    static_assert(!std::is_same_v<UInt, bool>);
    static_assert(0 < BitWidth);
    static_assert(BitWidth <= std::numeric_limits<UInt>::digits);

    using node_id = int;
    static constexpr node_id null_node = -1;

    struct Node {
        node_id child[2];
        int count;
        T prod;

        Node() : child{null_node, null_node}, count(0), prod(Monoid::id()) {}
    };

   private:
    struct Aggregate {
        int count;
        T prod;
    };

    std::vector<Node> nodes;
    UInt lazy_xor;

    static constexpr int bit(UInt value, int position) {
        return int((value >> position) & UInt(1));
    }

    static constexpr UInt value_mask() {
        if constexpr (BitWidth == std::numeric_limits<UInt>::digits) {
            return std::numeric_limits<UInt>::max();
        } else {
            return (UInt(1) << BitWidth) - UInt(1);
        }
    }

    static constexpr bool valid_value(UInt value) {
        return (value & ~value_mask()) == UInt(0);
    }

    node_id new_node() {
        nodes.emplace_back();
        return int(nodes.size()) - 1;
    }

    int subtree_size(node_id node) const {
        return node == null_node ? 0 : nodes[node].count;
    }

    T subtree_prod(node_id node) const {
        return node == null_node ? Monoid::id() : nodes[node].prod;
    }

    void update(int node) {
        nodes[node].count =
            subtree_size(nodes[node].child[0]) +
            subtree_size(nodes[node].child[1]);
        nodes[node].prod =
            Monoid::op(subtree_prod(nodes[node].child[0]),
                       subtree_prod(nodes[node].child[1]));
    }

    node_id find_node(UInt key) const {
        key ^= lazy_xor;
        node_id node = 0;
        for (int position = BitWidth - 1; position >= 0; --position) {
            node = nodes[node].child[bit(key, position)];
            if (node == null_node || nodes[node].count == 0) {
                return null_node;
            }
        }
        return node;
    }

    static int extend_comparison(int relation,
                                 int digit,
                                 int bound_digit) {
        if (relation != 0) return relation;
        if (digit < bound_digit) return -1;
        if (digit > bound_digit) return 1;
        return 0;
    }

    Aggregate xor_range_impl(int node,
                             int position,
                             UInt effective_xor,
                             UInt lower,
                             UInt upper,
                             int lower_relation,
                             int upper_relation) const {
        if (node == -1 || nodes[node].count == 0 ||
            lower_relation < 0 || upper_relation > 0) {
            return {0, Monoid::id()};
        }
        if (lower_relation > 0 && upper_relation < 0) {
            return {nodes[node].count, nodes[node].prod};
        }
        if (position < 0) {
            if (lower_relation >= 0 && upper_relation < 0) {
                return {nodes[node].count, nodes[node].prod};
            }
            return {0, Monoid::id()};
        }

        Aggregate result{0, Monoid::id()};
        const int xor_digit = bit(effective_xor, position);
        const int lower_digit = bit(lower, position);
        const int upper_digit = bit(upper, position);
        for (int xor_result_digit = 0;
             xor_result_digit < 2;
             ++xor_result_digit) {
            const int direction = xor_result_digit ^ xor_digit;
            Aggregate part = xor_range_impl(
                nodes[node].child[direction],
                position - 1,
                effective_xor,
                lower,
                upper,
                extend_comparison(lower_relation,
                                  xor_result_digit,
                                  lower_digit),
                extend_comparison(upper_relation,
                                  xor_result_digit,
                                  upper_digit));
            result.count += part.count;
            result.prod = Monoid::op(result.prod, part.prod);
        }
        return result;
    }

    T prod_xor_greater_equal_impl(UInt value, UInt lower) const {
        const UInt effective_xor = lazy_xor ^ value;
        T result = Monoid::id();
        int node = 0;
        for (int position = BitWidth - 1;
             position >= 0 && node != -1;
             --position) {
            const int zero = bit(effective_xor, position);
            if (bit(lower, position) == 0) {
                result =
                    Monoid::op(result,
                               subtree_prod(nodes[node].child[zero ^ 1]));
                node = nodes[node].child[zero];
            } else {
                node = nodes[node].child[zero ^ 1];
            }
        }
        return Monoid::op(result, subtree_prod(node));
    }

   public:
    BinaryTrieMonoid() : nodes(1), lazy_xor(0) {}

    BinaryTrieMonoid(
        std::initializer_list<std::pair<UInt, T>> init)
        : BinaryTrieMonoid() {
        for (const auto& entry : init) {
            insert(entry.first, entry.second);
        }
    }

    template <typename Iterator>
    BinaryTrieMonoid(Iterator first, Iterator last)
        : BinaryTrieMonoid() {
        while (first != last) {
            insert(first->first, first->second);
            ++first;
        }
    }

    BinaryTrieMonoid(const std::vector<UInt>& keys,
                     const std::vector<T>& values)
        : BinaryTrieMonoid() {
        assert(keys.size() == values.size());
        for (int i = 0; i < int(keys.size()); ++i) {
            insert(keys[i], values[i]);
        }
    }

    int size() const {
        return nodes[0].count;
    }

    bool empty() const {
        return size() == 0;
    }

    node_id root() const {
        return 0;
    }

    const Node& node(node_id id) const {
        assert(0 <= id && std::size_t(id) < nodes.size());
        return nodes[id];
    }

    node_id find(UInt key) const {
        assert(valid_value(key));
        return find_node(key);
    }

    std::size_t node_count() const {
        return nodes.size();
    }

    void reserve(std::size_t node_capacity) {
        nodes.reserve(node_capacity);
    }

    UInt xor_mask() const {
        return lazy_xor;
    }

    void clear() {
        nodes.clear();
        nodes.emplace_back();
        lazy_xor = 0;
    }

    node_id insert(UInt key, const T& value) {
        assert(valid_value(key));
        key ^= lazy_xor;
        node_id node = 0;
        ++nodes[node].count;
        nodes[node].prod = Monoid::op(nodes[node].prod, value);
        for (int position = BitWidth - 1; position >= 0; --position) {
            const int direction = bit(key, position);
            if (nodes[node].child[direction] == null_node) {
                const node_id child = new_node();
                nodes[node].child[direction] = child;
            }
            node = nodes[node].child[direction];
            ++nodes[node].count;
            nodes[node].prod = Monoid::op(nodes[node].prod, value);
        }
        return node;
    }

    int count(UInt key) const {
        assert(valid_value(key));
        const node_id node = find_node(key);
        return node == null_node ? 0 : nodes[node].count;
    }

    bool contains(UInt key) const {
        return count(key) > 0;
    }

    T prod(UInt key) const {
        assert(valid_value(key));
        const node_id node = find_node(key);
        return node == null_node ? Monoid::id() : nodes[node].prod;
    }

    T all_prod() const {
        return nodes[0].prod;
    }

    int erase_all(UInt key) {
        assert(valid_value(key));
        key ^= lazy_xor;

        int path[BitWidth + 1];
        path[0] = 0;
        int node = 0;
        for (int position = BitWidth - 1, depth = 1;
             position >= 0;
             --position, ++depth) {
            node = nodes[node].child[bit(key, position)];
            if (node == -1 || nodes[node].count == 0) return 0;
            path[depth] = node;
        }

        const int erased = nodes[node].count;
        nodes[node].count = 0;
        nodes[node].prod = Monoid::id();
        for (int depth = BitWidth - 1; depth >= 0; --depth) {
            update(path[depth]);
        }
        return erased;
    }

    void xor_all(UInt value) {
        assert(valid_value(value));
        lazy_xor ^= value;
    }

    UInt kth_xor(int k, UInt value) const {
        assert(0 <= k && k < size());
        assert(valid_value(value));
        const UInt effective_xor = lazy_xor ^ value;
        UInt result = 0;
        int node = 0;
        for (int position = BitWidth - 1; position >= 0; --position) {
            const int preferred = bit(effective_xor, position);
            const int preferred_size =
                subtree_size(nodes[node].child[preferred]);
            if (k < preferred_size) {
                node = nodes[node].child[preferred];
            } else {
                k -= preferred_size;
                node = nodes[node].child[preferred ^ 1];
                result |= UInt(1) << position;
            }
        }
        return result;
    }

    UInt kth(int k) const {
        return kth_xor(k, 0);
    }

    UInt min() const {
        return kth(0);
    }

    UInt max() const {
        return kth(size() - 1);
    }

    UInt min_xor(UInt value) const {
        return kth_xor(0, value);
    }

    UInt max_xor(UInt value) const {
        return kth_xor(size() - 1, value);
    }

    int count_xor_equal(UInt value, UInt target) const {
        assert(valid_value(value));
        assert(valid_value(target));
        return count(value ^ target);
    }

    int count_xor_less(UInt value, UInt upper) const {
        assert(valid_value(value));
        if (!valid_value(upper)) return size();

        const UInt effective_xor = lazy_xor ^ value;
        int result = 0;
        int node = 0;
        for (int position = BitWidth - 1;
             position >= 0 && node != -1;
             --position) {
            const int zero = bit(effective_xor, position);
            if (bit(upper, position) == 1) {
                result += subtree_size(nodes[node].child[zero]);
                node = nodes[node].child[zero ^ 1];
            } else {
                node = nodes[node].child[zero];
            }
        }
        return result;
    }

    int count_less_xor(UInt value, UInt upper) const {
        return count_xor_less(value, upper);
    }

    int count_xor_less_equal(UInt value, UInt upper) const {
        assert(valid_value(value));
        assert(valid_value(upper));
        if (upper == value_mask()) return size();
        return count_xor_less(value, upper + UInt(1));
    }

    int count_xor_greater(UInt value, UInt lower) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        return size() - count_xor_less_equal(value, lower);
    }

    int count_xor_greater_equal(UInt value, UInt lower) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        return size() - count_xor_less(value, lower);
    }

    int count_xor_range(UInt value, UInt lower, UInt upper) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        assert(lower <= upper);
        return count_xor_less(value, upper) -
               count_xor_less(value, lower);
    }

    int order_of_key(UInt key) const {
        return count_xor_less(0, key);
    }

    int count_less(UInt key) const {
        return order_of_key(key);
    }

    int count_less_equal(UInt key) const {
        return count_xor_less_equal(0, key);
    }

    int count_greater(UInt key) const {
        return count_xor_greater(0, key);
    }

    int count_greater_equal(UInt key) const {
        return count_xor_greater_equal(0, key);
    }

    int count_range(UInt lower, UInt upper) const {
        return count_xor_range(0, lower, upper);
    }

    T prod_xor_equal(UInt value, UInt target) const {
        assert(valid_value(value));
        assert(valid_value(target));
        return prod(value ^ target);
    }

    T prod_xor_less(UInt value, UInt upper) const {
        assert(valid_value(value));
        if (!valid_value(upper)) return all_prod();

        const UInt effective_xor = lazy_xor ^ value;
        T result = Monoid::id();
        int node = 0;
        for (int position = BitWidth - 1;
             position >= 0 && node != -1;
             --position) {
            const int zero = bit(effective_xor, position);
            if (bit(upper, position) == 1) {
                result =
                    Monoid::op(result,
                               subtree_prod(nodes[node].child[zero]));
                node = nodes[node].child[zero ^ 1];
            } else {
                node = nodes[node].child[zero];
            }
        }
        return result;
    }

    T prod_xor_less_equal(UInt value, UInt upper) const {
        assert(valid_value(value));
        assert(valid_value(upper));
        if (upper == value_mask()) return all_prod();
        return prod_xor_less(value, upper + UInt(1));
    }

    T prod_xor_greater(UInt value, UInt lower) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        if (lower == value_mask()) return Monoid::id();
        return prod_xor_greater_equal_impl(value, lower + UInt(1));
    }

    T prod_xor_greater_equal(UInt value, UInt lower) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        return prod_xor_greater_equal_impl(value, lower);
    }

    T prod_xor_range(UInt value, UInt lower, UInt upper) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        assert(lower <= upper);
        if (lower == upper) return Monoid::id();
        if (!valid_value(upper)) {
            return prod_xor_greater_equal(value, lower);
        }
        return xor_range_impl(0,
                              BitWidth - 1,
                              lazy_xor ^ value,
                              lower,
                              upper,
                              0,
                              0)
            .prod;
    }

    T prod_less(UInt key) const {
        return prod_xor_less(0, key);
    }

    T prod_less_equal(UInt key) const {
        return prod_xor_less_equal(0, key);
    }

    T prod_greater(UInt key) const {
        return prod_xor_greater(0, key);
    }

    T prod_greater_equal(UInt key) const {
        return prod_xor_greater_equal(0, key);
    }

    T prod_range(UInt lower, UInt upper) const {
        return prod_xor_range(0, lower, upper);
    }
};

}  // namespace ds
}  // namespace m1une

#endif  // M1UNE_DS_BINARY_TRIE_MONOID_HPP
#line 1 "ds/binary_trie/binary_trie_monoid.hpp"



#include <cassert>
#include <cstdint>
#include <initializer_list>
#include <limits>
#include <type_traits>
#include <utility>
#include <vector>

#line 1 "monoid/concept.hpp"



#include <concepts>

namespace m1une {
namespace monoid {

// Concept to check if a type satisfies the requirements of a Monoid.
// A Monoid must have a `value_type`, an identity element `id()`, and an associative binary operation `op()`.
template <typename M>
concept IsMonoid = requires(typename M::value_type a, typename M::value_type b) {
    // 1. Must define `value_type`
    typename M::value_type;

    // 2. Must have a static method `id()` returning `value_type`
    { M::id() } -> std::same_as<typename M::value_type>;

    // 3. Must have a static method `op(a, b)` returning `value_type`
    { M::op(a, b) } -> std::same_as<typename M::value_type>;
};

// Concept for groups. A type satisfying this concept must also obey the group
// laws; concepts can check the interface but not the algebraic properties.
template <typename M>
concept IsGroup = IsMonoid<M> && requires(typename M::value_type a) {
    { M::inv(a) } -> std::same_as<typename M::value_type>;
};

// Concept for commutative groups. Commutativity is a semantic requirement and
// cannot be checked by a C++ concept.
template <typename M>
concept IsCommutativeGroup = IsGroup<M>;

}  // namespace monoid
}  // namespace m1une


#line 13 "ds/binary_trie/binary_trie_monoid.hpp"

namespace m1une {
namespace ds {

template <m1une::monoid::IsMonoid Monoid,
          typename UInt = std::uint32_t,
          int BitWidth = std::numeric_limits<UInt>::digits>
struct BinaryTrieMonoid {
    using T = typename Monoid::value_type;

    static_assert(std::is_integral_v<UInt>);
    static_assert(std::is_unsigned_v<UInt>);
    static_assert(!std::is_same_v<UInt, bool>);
    static_assert(0 < BitWidth);
    static_assert(BitWidth <= std::numeric_limits<UInt>::digits);

    using node_id = int;
    static constexpr node_id null_node = -1;

    struct Node {
        node_id child[2];
        int count;
        T prod;

        Node() : child{null_node, null_node}, count(0), prod(Monoid::id()) {}
    };

   private:
    struct Aggregate {
        int count;
        T prod;
    };

    std::vector<Node> nodes;
    UInt lazy_xor;

    static constexpr int bit(UInt value, int position) {
        return int((value >> position) & UInt(1));
    }

    static constexpr UInt value_mask() {
        if constexpr (BitWidth == std::numeric_limits<UInt>::digits) {
            return std::numeric_limits<UInt>::max();
        } else {
            return (UInt(1) << BitWidth) - UInt(1);
        }
    }

    static constexpr bool valid_value(UInt value) {
        return (value & ~value_mask()) == UInt(0);
    }

    node_id new_node() {
        nodes.emplace_back();
        return int(nodes.size()) - 1;
    }

    int subtree_size(node_id node) const {
        return node == null_node ? 0 : nodes[node].count;
    }

    T subtree_prod(node_id node) const {
        return node == null_node ? Monoid::id() : nodes[node].prod;
    }

    void update(int node) {
        nodes[node].count =
            subtree_size(nodes[node].child[0]) +
            subtree_size(nodes[node].child[1]);
        nodes[node].prod =
            Monoid::op(subtree_prod(nodes[node].child[0]),
                       subtree_prod(nodes[node].child[1]));
    }

    node_id find_node(UInt key) const {
        key ^= lazy_xor;
        node_id node = 0;
        for (int position = BitWidth - 1; position >= 0; --position) {
            node = nodes[node].child[bit(key, position)];
            if (node == null_node || nodes[node].count == 0) {
                return null_node;
            }
        }
        return node;
    }

    static int extend_comparison(int relation,
                                 int digit,
                                 int bound_digit) {
        if (relation != 0) return relation;
        if (digit < bound_digit) return -1;
        if (digit > bound_digit) return 1;
        return 0;
    }

    Aggregate xor_range_impl(int node,
                             int position,
                             UInt effective_xor,
                             UInt lower,
                             UInt upper,
                             int lower_relation,
                             int upper_relation) const {
        if (node == -1 || nodes[node].count == 0 ||
            lower_relation < 0 || upper_relation > 0) {
            return {0, Monoid::id()};
        }
        if (lower_relation > 0 && upper_relation < 0) {
            return {nodes[node].count, nodes[node].prod};
        }
        if (position < 0) {
            if (lower_relation >= 0 && upper_relation < 0) {
                return {nodes[node].count, nodes[node].prod};
            }
            return {0, Monoid::id()};
        }

        Aggregate result{0, Monoid::id()};
        const int xor_digit = bit(effective_xor, position);
        const int lower_digit = bit(lower, position);
        const int upper_digit = bit(upper, position);
        for (int xor_result_digit = 0;
             xor_result_digit < 2;
             ++xor_result_digit) {
            const int direction = xor_result_digit ^ xor_digit;
            Aggregate part = xor_range_impl(
                nodes[node].child[direction],
                position - 1,
                effective_xor,
                lower,
                upper,
                extend_comparison(lower_relation,
                                  xor_result_digit,
                                  lower_digit),
                extend_comparison(upper_relation,
                                  xor_result_digit,
                                  upper_digit));
            result.count += part.count;
            result.prod = Monoid::op(result.prod, part.prod);
        }
        return result;
    }

    T prod_xor_greater_equal_impl(UInt value, UInt lower) const {
        const UInt effective_xor = lazy_xor ^ value;
        T result = Monoid::id();
        int node = 0;
        for (int position = BitWidth - 1;
             position >= 0 && node != -1;
             --position) {
            const int zero = bit(effective_xor, position);
            if (bit(lower, position) == 0) {
                result =
                    Monoid::op(result,
                               subtree_prod(nodes[node].child[zero ^ 1]));
                node = nodes[node].child[zero];
            } else {
                node = nodes[node].child[zero ^ 1];
            }
        }
        return Monoid::op(result, subtree_prod(node));
    }

   public:
    BinaryTrieMonoid() : nodes(1), lazy_xor(0) {}

    BinaryTrieMonoid(
        std::initializer_list<std::pair<UInt, T>> init)
        : BinaryTrieMonoid() {
        for (const auto& entry : init) {
            insert(entry.first, entry.second);
        }
    }

    template <typename Iterator>
    BinaryTrieMonoid(Iterator first, Iterator last)
        : BinaryTrieMonoid() {
        while (first != last) {
            insert(first->first, first->second);
            ++first;
        }
    }

    BinaryTrieMonoid(const std::vector<UInt>& keys,
                     const std::vector<T>& values)
        : BinaryTrieMonoid() {
        assert(keys.size() == values.size());
        for (int i = 0; i < int(keys.size()); ++i) {
            insert(keys[i], values[i]);
        }
    }

    int size() const {
        return nodes[0].count;
    }

    bool empty() const {
        return size() == 0;
    }

    node_id root() const {
        return 0;
    }

    const Node& node(node_id id) const {
        assert(0 <= id && std::size_t(id) < nodes.size());
        return nodes[id];
    }

    node_id find(UInt key) const {
        assert(valid_value(key));
        return find_node(key);
    }

    std::size_t node_count() const {
        return nodes.size();
    }

    void reserve(std::size_t node_capacity) {
        nodes.reserve(node_capacity);
    }

    UInt xor_mask() const {
        return lazy_xor;
    }

    void clear() {
        nodes.clear();
        nodes.emplace_back();
        lazy_xor = 0;
    }

    node_id insert(UInt key, const T& value) {
        assert(valid_value(key));
        key ^= lazy_xor;
        node_id node = 0;
        ++nodes[node].count;
        nodes[node].prod = Monoid::op(nodes[node].prod, value);
        for (int position = BitWidth - 1; position >= 0; --position) {
            const int direction = bit(key, position);
            if (nodes[node].child[direction] == null_node) {
                const node_id child = new_node();
                nodes[node].child[direction] = child;
            }
            node = nodes[node].child[direction];
            ++nodes[node].count;
            nodes[node].prod = Monoid::op(nodes[node].prod, value);
        }
        return node;
    }

    int count(UInt key) const {
        assert(valid_value(key));
        const node_id node = find_node(key);
        return node == null_node ? 0 : nodes[node].count;
    }

    bool contains(UInt key) const {
        return count(key) > 0;
    }

    T prod(UInt key) const {
        assert(valid_value(key));
        const node_id node = find_node(key);
        return node == null_node ? Monoid::id() : nodes[node].prod;
    }

    T all_prod() const {
        return nodes[0].prod;
    }

    int erase_all(UInt key) {
        assert(valid_value(key));
        key ^= lazy_xor;

        int path[BitWidth + 1];
        path[0] = 0;
        int node = 0;
        for (int position = BitWidth - 1, depth = 1;
             position >= 0;
             --position, ++depth) {
            node = nodes[node].child[bit(key, position)];
            if (node == -1 || nodes[node].count == 0) return 0;
            path[depth] = node;
        }

        const int erased = nodes[node].count;
        nodes[node].count = 0;
        nodes[node].prod = Monoid::id();
        for (int depth = BitWidth - 1; depth >= 0; --depth) {
            update(path[depth]);
        }
        return erased;
    }

    void xor_all(UInt value) {
        assert(valid_value(value));
        lazy_xor ^= value;
    }

    UInt kth_xor(int k, UInt value) const {
        assert(0 <= k && k < size());
        assert(valid_value(value));
        const UInt effective_xor = lazy_xor ^ value;
        UInt result = 0;
        int node = 0;
        for (int position = BitWidth - 1; position >= 0; --position) {
            const int preferred = bit(effective_xor, position);
            const int preferred_size =
                subtree_size(nodes[node].child[preferred]);
            if (k < preferred_size) {
                node = nodes[node].child[preferred];
            } else {
                k -= preferred_size;
                node = nodes[node].child[preferred ^ 1];
                result |= UInt(1) << position;
            }
        }
        return result;
    }

    UInt kth(int k) const {
        return kth_xor(k, 0);
    }

    UInt min() const {
        return kth(0);
    }

    UInt max() const {
        return kth(size() - 1);
    }

    UInt min_xor(UInt value) const {
        return kth_xor(0, value);
    }

    UInt max_xor(UInt value) const {
        return kth_xor(size() - 1, value);
    }

    int count_xor_equal(UInt value, UInt target) const {
        assert(valid_value(value));
        assert(valid_value(target));
        return count(value ^ target);
    }

    int count_xor_less(UInt value, UInt upper) const {
        assert(valid_value(value));
        if (!valid_value(upper)) return size();

        const UInt effective_xor = lazy_xor ^ value;
        int result = 0;
        int node = 0;
        for (int position = BitWidth - 1;
             position >= 0 && node != -1;
             --position) {
            const int zero = bit(effective_xor, position);
            if (bit(upper, position) == 1) {
                result += subtree_size(nodes[node].child[zero]);
                node = nodes[node].child[zero ^ 1];
            } else {
                node = nodes[node].child[zero];
            }
        }
        return result;
    }

    int count_less_xor(UInt value, UInt upper) const {
        return count_xor_less(value, upper);
    }

    int count_xor_less_equal(UInt value, UInt upper) const {
        assert(valid_value(value));
        assert(valid_value(upper));
        if (upper == value_mask()) return size();
        return count_xor_less(value, upper + UInt(1));
    }

    int count_xor_greater(UInt value, UInt lower) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        return size() - count_xor_less_equal(value, lower);
    }

    int count_xor_greater_equal(UInt value, UInt lower) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        return size() - count_xor_less(value, lower);
    }

    int count_xor_range(UInt value, UInt lower, UInt upper) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        assert(lower <= upper);
        return count_xor_less(value, upper) -
               count_xor_less(value, lower);
    }

    int order_of_key(UInt key) const {
        return count_xor_less(0, key);
    }

    int count_less(UInt key) const {
        return order_of_key(key);
    }

    int count_less_equal(UInt key) const {
        return count_xor_less_equal(0, key);
    }

    int count_greater(UInt key) const {
        return count_xor_greater(0, key);
    }

    int count_greater_equal(UInt key) const {
        return count_xor_greater_equal(0, key);
    }

    int count_range(UInt lower, UInt upper) const {
        return count_xor_range(0, lower, upper);
    }

    T prod_xor_equal(UInt value, UInt target) const {
        assert(valid_value(value));
        assert(valid_value(target));
        return prod(value ^ target);
    }

    T prod_xor_less(UInt value, UInt upper) const {
        assert(valid_value(value));
        if (!valid_value(upper)) return all_prod();

        const UInt effective_xor = lazy_xor ^ value;
        T result = Monoid::id();
        int node = 0;
        for (int position = BitWidth - 1;
             position >= 0 && node != -1;
             --position) {
            const int zero = bit(effective_xor, position);
            if (bit(upper, position) == 1) {
                result =
                    Monoid::op(result,
                               subtree_prod(nodes[node].child[zero]));
                node = nodes[node].child[zero ^ 1];
            } else {
                node = nodes[node].child[zero];
            }
        }
        return result;
    }

    T prod_xor_less_equal(UInt value, UInt upper) const {
        assert(valid_value(value));
        assert(valid_value(upper));
        if (upper == value_mask()) return all_prod();
        return prod_xor_less(value, upper + UInt(1));
    }

    T prod_xor_greater(UInt value, UInt lower) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        if (lower == value_mask()) return Monoid::id();
        return prod_xor_greater_equal_impl(value, lower + UInt(1));
    }

    T prod_xor_greater_equal(UInt value, UInt lower) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        return prod_xor_greater_equal_impl(value, lower);
    }

    T prod_xor_range(UInt value, UInt lower, UInt upper) const {
        assert(valid_value(value));
        assert(valid_value(lower));
        assert(lower <= upper);
        if (lower == upper) return Monoid::id();
        if (!valid_value(upper)) {
            return prod_xor_greater_equal(value, lower);
        }
        return xor_range_impl(0,
                              BitWidth - 1,
                              lazy_xor ^ value,
                              lower,
                              upper,
                              0,
                              0)
            .prod;
    }

    T prod_less(UInt key) const {
        return prod_xor_less(0, key);
    }

    T prod_less_equal(UInt key) const {
        return prod_xor_less_equal(0, key);
    }

    T prod_greater(UInt key) const {
        return prod_xor_greater(0, key);
    }

    T prod_greater_equal(UInt key) const {
        return prod_xor_greater_equal(0, key);
    }

    T prod_range(UInt lower, UInt upper) const {
        return prod_xor_range(0, lower, upper);
    }
};

}  // namespace ds
}  // namespace m1une
Back to top page