This documentation is automatically generated by online-judge-tools/verification-helper
#include "ds/binary_trie.hpp"#include "other/bit.hpp"
#include "ds/node_pool.hpp"
// 非永続ならば、2 * 要素数 のノード数
template <int LOG, bool PERSISTENT, typename UINT = u64,
typename SIZE_TYPE = u32>
struct Binary_Trie {
using T = SIZE_TYPE;
static_assert(is_same_v<T, u32> || is_same_v<T, u64>);
static_assert(0 < LOG && LOG <= numeric_limits<UINT>::digits);
struct Node {
int width;
UINT val;
T cnt;
Node *l, *r;
};
Node_Pool<Node> pool;
using np = Node *;
void reset() { pool.reset(); }
np new_root() { return nullptr; }
np add(np root, UINT val, T cnt = 1) {
if (!root) root = new_node(0, 0);
assert((val >> LOG) == 0);
return add_rec(root, LOG, val, cnt);
}
// f(val, cnt)
template <typename F>
void enumerate(np root, F f) {
auto dfs = [&](auto &dfs, np root, UINT val, int ht) -> void {
if (ht == 0) {
f(val, root->cnt);
return;
}
np c = root->l;
if (c) {
dfs(dfs, c, val << (c->width) | (c->val), ht - (c->width));
}
c = root->r;
if (c) {
dfs(dfs, c, val << (c->width) | (c->val), ht - (c->width));
}
};
if (root) dfs(dfs, root, 0, LOG);
}
// xor_val したあとの値で昇順 k 番目
UINT kth(np root, T k, UINT xor_val) {
assert(root && k < root->cnt);
return kth_rec(root, 0, k, LOG, xor_val) ^ xor_val;
}
// xor_val したあとの値で最小値
UINT min(np root, UINT xor_val) {
assert(root && root->cnt);
return kth(root, 0, xor_val);
}
// xor_val したあとの値で最大値
UINT max(np root, UINT xor_val) {
assert(root && root->cnt);
return kth(root, (root->cnt) - 1, xor_val);
}
// xor_val したあとの値で [0, upper) 内に入るものの個数
T prefix_count(np root, UINT upper, UINT xor_val) {
if (!root) return 0;
return prefix_count_rec(root, LOG, upper, xor_val, 0);
}
// xor_val したあとの値で [lo, hi) 内に入るものの個数
T count(np root, UINT lo, UINT hi, UINT xor_val) {
return prefix_count(root, hi, xor_val) - prefix_count(root, lo, xor_val);
}
private:
inline UINT mask(int k) { return (UINT(1) << k) - 1; }
np new_node(int width, UINT val) {
np c = pool.create();
c->l = c->r = nullptr;
c->width = width, c->val = val, c->cnt = 0;
return c;
}
np clone(np c) {
if (!c || !PERSISTENT) return c;
return pool.clone(c);
}
np add_rec(np root, int ht, UINT val, T cnt) {
root = clone(root);
root->cnt += cnt;
if (ht == 0) return root;
bool go_r = (val >> (ht - 1)) & 1;
np c = (go_r ? root->r : root->l);
if (!c) {
c = new_node(ht, val);
c->cnt = cnt;
if (!go_r) root->l = c;
if (go_r) root->r = c;
return root;
}
int w = c->width;
if ((val >> (ht - w)) == c->val) {
c = add_rec(c, ht - w, val & mask(ht - w), cnt);
if (!go_r) root->l = c;
if (go_r) root->r = c;
return root;
}
int same = w - 1 - topbit((val >> (ht - w)) ^ (c->val));
np n = new_node(same, (c->val) >> (w - same));
n->cnt = c->cnt + cnt;
c = clone(c);
c->width = w - same;
c->val = c->val & mask(w - same);
if ((val >> (ht - same - 1)) & 1) {
n->l = c;
n->r = new_node(ht - same, val & mask(ht - same));
n->r->cnt = cnt;
} else {
n->r = c;
n->l = new_node(ht - same, val & mask(ht - same));
n->l->cnt = cnt;
}
if (!go_r) root->l = n;
if (go_r) root->r = n;
return root;
}
UINT kth_rec(np root, UINT val, T k, int ht, UINT xor_val) {
if (ht == 0) return val;
np left = root->l, right = root->r;
if ((xor_val >> (ht - 1)) & 1) swap(left, right);
T sl = (left ? left->cnt : 0);
np c;
if (k < sl) {
c = left;
}
if (k >= sl) {
c = right, k -= sl;
}
int w = c->width;
return kth_rec(c, val << w | (c->val), k, ht - w, xor_val);
}
T prefix_count_rec(np root, int ht, UINT LIM, UINT xor_val, UINT val) {
UINT now = (val << ht) ^ (xor_val);
if ((LIM >> ht) > (now >> ht)) return root->cnt;
if (ht == 0 || (LIM >> ht) < (now >> ht)) return 0;
T res = 0;
FOR(k, 2) {
np c = (k == 0 ? root->l : root->r);
if (c) {
int w = c->width;
res += prefix_count_rec(c, ht - w, LIM, xor_val, val << w | c->val);
}
}
return res;
}
};#line 1 "other/bit.hpp"
int popcnt(int x) { return __builtin_popcount(x); }
int popcnt(u32 x) { return __builtin_popcount(x); }
int popcnt(ll x) { return __builtin_popcountll(x); }
int popcnt(u64 x) { return __builtin_popcountll(x); }
int popcnt_sgn(int x) { return (__builtin_parity(unsigned(x)) & 1 ? -1 : 1); }
int popcnt_sgn(u32 x) { return (__builtin_parity(x) & 1 ? -1 : 1); }
int popcnt_sgn(ll x) { return (__builtin_parityll(x) & 1 ? -1 : 1); }
int popcnt_sgn(u64 x) { return (__builtin_parityll(x) & 1 ? -1 : 1); }
// (0, 1, 2, 3, 4) -> (-1, 0, 1, 1, 2)
int topbit(int x) { return (x == 0 ? -1 : 31 - __builtin_clz(x)); }
int topbit(u32 x) { return (x == 0 ? -1 : 31 - __builtin_clz(x)); }
int topbit(ll x) { return (x == 0 ? -1 : 63 - __builtin_clzll(x)); }
int topbit(u64 x) { return (x == 0 ? -1 : 63 - __builtin_clzll(x)); }
// (0, 1, 2, 3, 4) -> (-1, 0, 1, 0, 2)
int lowbit(int x) { return (x == 0 ? -1 : __builtin_ctz(x)); }
int lowbit(u32 x) { return (x == 0 ? -1 : __builtin_ctz(x)); }
int lowbit(ll x) { return (x == 0 ? -1 : __builtin_ctzll(x)); }
int lowbit(u64 x) { return (x == 0 ? -1 : __builtin_ctzll(x)); }
template <typename T>
T kth_bit(int k) {
assert(0 <= k && k < int(8 * sizeof(T)));
return T(1) << k;
}
template <typename T>
bool has_kth_bit(T x, int k) {
assert(0 <= k && k < int(8 * sizeof(T)));
return x >> k & 1;
}
template <typename UINT>
struct all_bit {
static_assert(is_unsigned<UINT>::value);
UINT s;
all_bit(UINT s) : s(s) {}
struct iter {
UINT s;
int operator*() const { return lowbit(s); }
void operator++() { s &= s - 1; }
bool operator!=(nullptr_t) const { return s; }
};
iter begin() const { return {s}; }
nullptr_t end() const { return nullptr; }
};
template <typename UINT>
struct all_subset {
static_assert(is_unsigned<UINT>::value);
UINT s;
all_subset(UINT s) : s(s) {}
struct iter {
UINT s, t;
bool done = false;
UINT operator*() const { return t; }
void operator++() {
done = (t == 0);
t = (t - 1) & s;
}
bool operator!=(nullptr_t) const { return !done; }
};
iter begin() const { return {s, s}; }
nullptr_t end() const { return nullptr; }
};
constexpr u64 full_mask(int n) {
assert(0 <= n && n <= 64);
return n == 64 ? -1ULL : (1ULL << n) - 1;
}
u64 bit_reverse(u64 x) {
x = ((x & 0x5555555555555555ULL) << 1) | ((x >> 1) & 0x5555555555555555ULL);
x = ((x & 0x3333333333333333ULL) << 2) | ((x >> 2) & 0x3333333333333333ULL);
x = ((x & 0x0f0f0f0f0f0f0f0fULL) << 4) | ((x >> 4) & 0x0f0f0f0f0f0f0f0fULL);
x = ((x & 0x00ff00ff00ff00ffULL) << 8) | ((x >> 8) & 0x00ff00ff00ff00ffULL);
x = ((x & 0x0000ffff0000ffffULL) << 16) | ((x >> 16) & 0x0000ffff0000ffffULL);
x = (x << 32) | (x >> 32);
return x;
}
#line 1 "ds/node_pool.hpp"
// マルチテストケースでも確保済み chunk を再利用する
template <class Node>
struct Node_Pool {
union Slot {
Node node;
Slot* next;
Slot() {}
~Slot() {}
};
using np = Node*;
static constexpr int CHUNK_SIZE = 1 << 12;
vc<unique_ptr<Slot[]>> chunks;
int chunk_id = 0;
int pos = 0;
Slot* free_head = nullptr;
~Node_Pool() {
auto& cache = chunk_cache();
for (auto& p : chunks) cache.eb(std::move(p));
}
template <class... Args>
np create(Args&&... args) {
Slot* s = new_slot();
return ::new (&s->node) Node(forward<Args>(args)...);
}
np clone(const np x) {
assert(x);
Slot* s = new_slot();
return ::new (&s->node) Node(*x);
}
void destroy(np x) {
if (!x) return;
x->~Node();
Slot* s = reinterpret_cast<Slot*>(x);
s->next = free_head;
free_head = s;
}
// 全 node を無効化する。
// 確保済み chunk は解放せず、次回以降に再利用する。
void reset() {
free_head = nullptr;
chunk_id = 0;
pos = 0;
}
private:
static vc<unique_ptr<Slot[]>>& chunk_cache() {
// static Node_Pool の destructor より先に破棄されないようにする。
static auto* cache = new vc<unique_ptr<Slot[]>>();
return *cache;
}
void alloc_chunk() {
auto& cache = chunk_cache();
if (cache.empty()) {
chunks.eb(make_unique<Slot[]>(CHUNK_SIZE));
} else {
chunks.eb(std::move(cache.back()));
cache.pop_back();
}
}
Slot* new_slot() {
if (free_head) {
Slot* s = free_head;
free_head = free_head->next;
return s;
}
if (chunk_id == len(chunks)) alloc_chunk();
Slot* s = &chunks[chunk_id][pos++];
if (pos == CHUNK_SIZE) {
++chunk_id;
pos = 0;
}
return s;
}
};
#line 3 "ds/binary_trie.hpp"
// 非永続ならば、2 * 要素数 のノード数
template <int LOG, bool PERSISTENT, typename UINT = u64,
typename SIZE_TYPE = u32>
struct Binary_Trie {
using T = SIZE_TYPE;
static_assert(is_same_v<T, u32> || is_same_v<T, u64>);
static_assert(0 < LOG && LOG <= numeric_limits<UINT>::digits);
struct Node {
int width;
UINT val;
T cnt;
Node *l, *r;
};
Node_Pool<Node> pool;
using np = Node *;
void reset() { pool.reset(); }
np new_root() { return nullptr; }
np add(np root, UINT val, T cnt = 1) {
if (!root) root = new_node(0, 0);
assert((val >> LOG) == 0);
return add_rec(root, LOG, val, cnt);
}
// f(val, cnt)
template <typename F>
void enumerate(np root, F f) {
auto dfs = [&](auto &dfs, np root, UINT val, int ht) -> void {
if (ht == 0) {
f(val, root->cnt);
return;
}
np c = root->l;
if (c) {
dfs(dfs, c, val << (c->width) | (c->val), ht - (c->width));
}
c = root->r;
if (c) {
dfs(dfs, c, val << (c->width) | (c->val), ht - (c->width));
}
};
if (root) dfs(dfs, root, 0, LOG);
}
// xor_val したあとの値で昇順 k 番目
UINT kth(np root, T k, UINT xor_val) {
assert(root && k < root->cnt);
return kth_rec(root, 0, k, LOG, xor_val) ^ xor_val;
}
// xor_val したあとの値で最小値
UINT min(np root, UINT xor_val) {
assert(root && root->cnt);
return kth(root, 0, xor_val);
}
// xor_val したあとの値で最大値
UINT max(np root, UINT xor_val) {
assert(root && root->cnt);
return kth(root, (root->cnt) - 1, xor_val);
}
// xor_val したあとの値で [0, upper) 内に入るものの個数
T prefix_count(np root, UINT upper, UINT xor_val) {
if (!root) return 0;
return prefix_count_rec(root, LOG, upper, xor_val, 0);
}
// xor_val したあとの値で [lo, hi) 内に入るものの個数
T count(np root, UINT lo, UINT hi, UINT xor_val) {
return prefix_count(root, hi, xor_val) - prefix_count(root, lo, xor_val);
}
private:
inline UINT mask(int k) { return (UINT(1) << k) - 1; }
np new_node(int width, UINT val) {
np c = pool.create();
c->l = c->r = nullptr;
c->width = width, c->val = val, c->cnt = 0;
return c;
}
np clone(np c) {
if (!c || !PERSISTENT) return c;
return pool.clone(c);
}
np add_rec(np root, int ht, UINT val, T cnt) {
root = clone(root);
root->cnt += cnt;
if (ht == 0) return root;
bool go_r = (val >> (ht - 1)) & 1;
np c = (go_r ? root->r : root->l);
if (!c) {
c = new_node(ht, val);
c->cnt = cnt;
if (!go_r) root->l = c;
if (go_r) root->r = c;
return root;
}
int w = c->width;
if ((val >> (ht - w)) == c->val) {
c = add_rec(c, ht - w, val & mask(ht - w), cnt);
if (!go_r) root->l = c;
if (go_r) root->r = c;
return root;
}
int same = w - 1 - topbit((val >> (ht - w)) ^ (c->val));
np n = new_node(same, (c->val) >> (w - same));
n->cnt = c->cnt + cnt;
c = clone(c);
c->width = w - same;
c->val = c->val & mask(w - same);
if ((val >> (ht - same - 1)) & 1) {
n->l = c;
n->r = new_node(ht - same, val & mask(ht - same));
n->r->cnt = cnt;
} else {
n->r = c;
n->l = new_node(ht - same, val & mask(ht - same));
n->l->cnt = cnt;
}
if (!go_r) root->l = n;
if (go_r) root->r = n;
return root;
}
UINT kth_rec(np root, UINT val, T k, int ht, UINT xor_val) {
if (ht == 0) return val;
np left = root->l, right = root->r;
if ((xor_val >> (ht - 1)) & 1) swap(left, right);
T sl = (left ? left->cnt : 0);
np c;
if (k < sl) {
c = left;
}
if (k >= sl) {
c = right, k -= sl;
}
int w = c->width;
return kth_rec(c, val << w | (c->val), k, ht - w, xor_val);
}
T prefix_count_rec(np root, int ht, UINT LIM, UINT xor_val, UINT val) {
UINT now = (val << ht) ^ (xor_val);
if ((LIM >> ht) > (now >> ht)) return root->cnt;
if (ht == 0 || (LIM >> ht) < (now >> ht)) return 0;
T res = 0;
FOR(k, 2) {
np c = (k == 0 ? root->l : root->r);
if (c) {
int w = c->width;
res += prefix_count_rec(c, ht - w, LIM, xor_val, val << w | c->val);
}
}
return res;
}
};