Binary Trie with Monoid
(ds/binary_trie/binary_trie_monoid.hpp)
- View this file on GitHub
- Last update: 2026-07-16 20:44:42+09:00
- Include:
#include "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
-
Monoid: A monoid satisfyingm1une::monoid::IsMonoid. Its operation must also be commutative. -
UInt: An unsigned integer type used for keys. Defaults tostd::uint32_t. -
BitWidth: The number of low key bits used by the trie. Defaults to all bits ofUInt. Keys and xor operands must fit in these bits.
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