20 using node_index_t = size_t;
21 static constexpr node_index_t npos = std::numeric_limits<node_index_t>::max();
23 struct PreWaveletNode {
24 node_index_t parent = npos;
25 node_index_t left_child = npos;
26 node_index_t right_child = npos;
29 explicit PreWaveletNode(uint64_t middle) : middle(middle) {}
41 node_index_t parent, left_child, right_child;
43 Storage bit_vector_data;
50 std::copy(data.begin(), data.end(), view.begin());
51 return std::move(result);
54 WaveletNode() =
default;
56 WaveletNode(PreWaveletNode&& node)
57 requires(std::same_as<Storage, AlignedStorage>)
58 : parent(node.parent),
59 left_child(node.left_child),
60 right_child(node.right_child),
62 bit_vector_data(std::move(align(node.stream.extract()))),
63 data(bit_vector_data.as_words64(), node.stream.size()) {}
67 bs << parent << left_child << right_child << middle;
68 bit_vector_data.serialize(bs);
73 static WaveletNode deserialize(std::span<const std::byte>& data)
74 requires(std::same_as<Storage, ReadOnlyStorageView>)
77 auto read = [&data](
auto& value) {
78 constexpr size_t length =
sizeof(value);
79 std::memcpy(&value, data.data(), length);
80 data = data.subspan(length);
83 read(result.left_child);
84 read(result.right_child);
87 result.data = RankSelectSupport<Storage>::deserialize(
88 result.bit_vector_data.as_words64(), data);
93 size_t alphabet_size_, data_size_;
95 std::vector<WaveletNode> nodes_;
96 std::vector<node_index_t> leaves_;
97 std::vector<size_t> permutation_, inverse_permutation_;
114 template <
typename F>
115 node_index_t build_node(
size_t begin,
119 std::span<const size_t> prefix_sum,
120 std::vector<PreWaveletNode>& nodes)
121 requires(std::same_as<Storage, AlignedStorage>)
123 if (end - begin == 1) {
124 leaves_[begin] = parent;
127 if (prefix_sum[end] == prefix_sum[begin]) {
128 for (
size_t symbol = begin; symbol < end; symbol++) {
129 leaves_[symbol] = parent;
134 node_index_t result = nodes.size();
135 size_t middle = get_middle(result);
136 middle = begin + (middle == npos ? (end - begin) / 2 : middle);
138 nodes.emplace_back(middle);
139 nodes[result].stream.reserve(prefix_sum[end] - prefix_sum[begin]);
140 nodes[result].parent = parent;
141 nodes[result].left_child =
142 build_node(begin, middle, result, get_middle, prefix_sum, nodes);
143 nodes[result].right_child =
144 build_node(middle, end, result, get_middle, prefix_sum, nodes);
164 void copy_segment_content(node_index_t node,
167 std::span<uint64_t> dst,
168 std::span<uint64_t> tmp)
const {
172 const size_t rank = nodes_[node].data.rank(begin), rank0 = begin -
rank;
173 const size_t right = nodes_[node].data.rank(end) -
rank,
174 left = (end - begin) - right;
176 if (nodes_[node].left_child == npos) {
177 std::fill_n(tmp.begin(),
static_cast<long long>(left),
178 inverse_permutation_[nodes_[node].middle - 1]);
180 copy_segment_content(nodes_[node].left_child, rank0, rank0 + left,
181 tmp.subspan(0, left), dst.subspan(0, left));
183 if (nodes_[node].right_child == npos) {
184 std::fill(tmp.begin() +
static_cast<long long>(left), tmp.end(),
185 inverse_permutation_[nodes_[node].middle]);
187 copy_segment_content(nodes_[node].right_child,
rank,
rank + right,
188 tmp.subspan(left, right), dst.subspan(left, right));
191 size_t j = 0, k = left;
192 const auto& bit_vector = nodes_[node].bit_vector_data.as_words64();
193 for (
size_t i = begin; i < end; i++) {
194 if ((bit_vector[i / 64] >> (i % 64)) & 1) {
195 dst[i - begin] = tmp[k++];
197 dst[i - begin] = tmp[j++];
202 WaveletTreeIndex() =
default;
220 size_t alphabet_size,
221 std::span<const uint64_t> data,
223 requires(std::same_as<Storage, AlignedStorage>)
224 : alphabet_size_(alphabet_size),
225 data_size_(data.
size()),
226 leaves_(alphabet_size_, npos) {
227 if (alphabet_size == 0) {
231 std::vector<PreWaveletNode> nodes;
232 nodes.reserve(alphabet_size_);
233 std::vector<size_t> nodes_structure;
234 nodes_structure.reserve(alphabet_size_);
236 if (build_type == WaveletTreeBuildType::Standard) {
237 permutation_.resize(alphabet_size);
238 inverse_permutation_.resize(alphabet_size);
239 std::iota(permutation_.begin(), permutation_.end(), 0);
240 std::iota(inverse_permutation_.begin(), inverse_permutation_.end(), 0);
241 nodes_structure.resize(alphabet_size_, npos);
244 size_t size, left, right;
246 std::vector<Node> huffman_nodes(alphabet_size_, {0, 0, 0});
247 for (
auto symb : data) {
248 huffman_nodes[symb].size++;
251 using elem_t = std::pair<size_t, size_t>;
252 std::priority_queue<elem_t, std::vector<elem_t>, std::greater<>> queue;
253 for (
size_t i = 0; i < alphabet_size_; i++) {
254 queue.emplace(huffman_nodes[i].
size, i);
256 while (queue.size() >= 2) {
257 auto right = queue.top().second;
259 auto left = queue.top().second;
261 huffman_nodes.push_back(
262 {huffman_nodes[left].size + huffman_nodes[right].size, left + 1,
264 queue.emplace(huffman_nodes.back().size, huffman_nodes.size() - 1);
267 std::function<size_t(
size_t)> enumerate = [&](
size_t index) ->
size_t {
268 const auto& [
size, left, right] = huffman_nodes[index];
269 if (left == 0 || right == 0) {
270 permutation_[index] = inverse_permutation_.size();
271 inverse_permutation_.push_back(index);
274 size_t ind = nodes_structure.size(), subtree = 0;
276 nodes_structure.push_back(0);
278 subtree += enumerate(left - 1);
280 nodes_structure[ind] = subtree;
282 subtree += enumerate(right - 1);
286 permutation_.resize(alphabet_size_);
287 inverse_permutation_.reserve(alphabet_size_);
288 enumerate(huffman_nodes.size() - 1);
291 std::vector<size_t> prefix_sum(alphabet_size + 1);
292 for (
auto symbol : data) {
293 prefix_sum[permutation_[symbol] + 1]++;
295 for (
size_t i = 0; i < alphabet_size_; i++) {
296 prefix_sum[i + 1] += prefix_sum[i];
300 0, alphabet_size_, npos,
301 [&](node_index_t node) {
return nodes_structure[node]; }, prefix_sum,
303 for (
auto symbol : data) {
304 auto index = permutation_[symbol];
305 for (node_index_t current = root_; current != npos;) {
306 auto& node = nodes[current];
307 bool go_right = index >= node.middle;
308 node.stream << go_right;
310 current = node.right_child;
312 current = node.left_child;
316 nodes_.reserve(nodes.size());
317 for (
auto& node : nodes) {
318 nodes_.emplace_back(std::move(node));
331 if (symbol >= alphabet_size_) [[unlikely]] {
334 symbol = permutation_[symbol];
335 for (node_index_t current = root_; current != npos;) {
336 const WaveletNode& node = nodes_[current];
337 if (symbol < node.middle) {
338 pos = node.data.
rank0(pos);
339 current = node.left_child;
341 pos = node.data.
rank(pos);
342 current = node.right_child;
357 if (symbol >= alphabet_size_ || data_size_ == 0) [[unlikely]] {
360 symbol = permutation_[symbol];
361 node_index_t current = leaves_[symbol];
362 for (; current != npos; current = nodes_[current].parent) {
363 const WaveletNode& node = nodes_[current];
364 if (symbol < node.middle) {
382 if (alphabet_size_ == 0 || data_size_ == 0 || begin >= end) [[unlikely]] {
385 auto length =
static_cast<long long>(end - begin);
386 std::vector<uint64_t> result(2 * length);
387 copy_segment_content(root_, begin, end,
388 std::span{result.begin(), result.begin() + length},
389 std::span{result.begin() + length, result.end()});
390 result.resize(length);
405 bs << alphabet_size_ << data_size_ << root_ << nodes_.
size();
406 for (
const WaveletNode& node : nodes_) {
409 for (
const node_index_t leaf : leaves_) {
412 for (
const size_t idx : permutation_) {
418 std::span<const std::byte>& data) {
420 auto read = [&data](
auto& value) {
421 constexpr size_t length =
sizeof(value);
422 std::memcpy(&value, data.data(), length);
423 data = data.subspan(length);
425 read(result.alphabet_size_);
426 read(result.data_size_);
430 result.nodes_.resize(
size);
431 for (
auto& node : result.nodes_) {
432 node = WaveletNode::deserialize(data);
434 result.leaves_.resize(result.alphabet_size_);
435 for (node_index_t& leaf : result.leaves_) {
438 result.permutation_.resize(result.alphabet_size_);
439 for (
size_t& index : result.permutation_) {
442 result.inverse_permutation_.resize(result.alphabet_size_);
443 for (
size_t i = 0; i < result.alphabet_size_; i++) {
444 result.inverse_permutation_[result.permutation_[i]] = i;