Pixie
Loading...
Searching...
No Matches
index.h
1#pragma once
2
3#include <pixie/rank_select/support.h>
5
6#include <cstdlib>
7#include <cstring>
8#include <functional>
9#include <limits>
10#include <numeric>
11#include <queue>
12#include <span>
13#include <vector>
14
15namespace pixie {
16
17template <StorageImplementation Storage = AlignedStorage>
18class WaveletTreeIndex : public WaveletTreeBase<WaveletTreeIndex<Storage>> {
19 private:
20 using node_index_t = size_t;
21 static constexpr node_index_t npos = std::numeric_limits<node_index_t>::max();
22
23 struct PreWaveletNode {
24 node_index_t parent = npos;
25 node_index_t left_child = npos;
26 node_index_t right_child = npos;
27 uint64_t middle;
28 OutputBitStream stream;
29 explicit PreWaveletNode(uint64_t middle) : middle(middle) {}
30 };
31
40 struct WaveletNode {
41 node_index_t parent, left_child, right_child;
42 uint64_t middle;
43 Storage bit_vector_data;
45
47 static AlignedStorage align(std::vector<uint64_t>&& data) {
48 AlignedStorage result(data.size() * 64);
49 auto view = result.writable_words64();
50 std::copy(data.begin(), data.end(), view.begin());
51 return std::move(result);
52 }
53
54 WaveletNode() = default;
55
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),
61 middle(node.middle),
62 bit_vector_data(std::move(align(node.stream.extract()))),
63 data(bit_vector_data.as_words64(), node.stream.size()) {}
64
66 void serialize(pixie::OutputBitStream& bs) const {
67 bs << parent << left_child << right_child << middle;
68 bit_vector_data.serialize(bs);
69 data.serialize(bs);
70 }
71
73 static WaveletNode deserialize(std::span<const std::byte>& data)
74 requires(std::same_as<Storage, ReadOnlyStorageView>)
75 {
76 WaveletNode result;
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);
81 };
82 read(result.parent);
83 read(result.left_child);
84 read(result.right_child);
85 read(result.middle);
86 result.bit_vector_data = ReadOnlyStorageView::deserialize(data);
87 result.data = RankSelectSupport<Storage>::deserialize(
88 result.bit_vector_data.as_words64(), data);
89 return result;
90 }
91 };
92
93 size_t alphabet_size_, data_size_;
94 node_index_t root_;
95 std::vector<WaveletNode> nodes_;
96 std::vector<node_index_t> leaves_;
97 std::vector<size_t> permutation_, inverse_permutation_;
98
114 template <typename F>
115 node_index_t build_node(size_t begin,
116 size_t end,
117 node_index_t parent,
118 const F& get_middle,
119 std::span<const size_t> prefix_sum,
120 std::vector<PreWaveletNode>& nodes)
121 requires(std::same_as<Storage, AlignedStorage>)
122 {
123 if (end - begin == 1) {
124 leaves_[begin] = parent;
125 return npos;
126 }
127 if (prefix_sum[end] == prefix_sum[begin]) {
128 for (size_t symbol = begin; symbol < end; symbol++) {
129 leaves_[symbol] = parent;
130 }
131 return npos;
132 }
133
134 node_index_t result = nodes.size();
135 size_t middle = get_middle(result);
136 middle = begin + (middle == npos ? (end - begin) / 2 : middle);
137
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);
145
146 return result;
147 }
148
164 void copy_segment_content(node_index_t node,
165 size_t begin,
166 size_t end,
167 std::span<uint64_t> dst,
168 std::span<uint64_t> tmp) const {
169 if (begin == end) {
170 return;
171 }
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;
175
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]);
179 } else {
180 copy_segment_content(nodes_[node].left_child, rank0, rank0 + left,
181 tmp.subspan(0, left), dst.subspan(0, left));
182 }
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]);
186 } else {
187 copy_segment_content(nodes_[node].right_child, rank, rank + right,
188 tmp.subspan(left, right), dst.subspan(left, right));
189 }
190
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++];
196 } else {
197 dst[i - begin] = tmp[j++];
198 }
199 }
200 }
201
202 WaveletTreeIndex() = default;
203
204 public:
220 size_t alphabet_size,
221 std::span<const uint64_t> data,
222 const WaveletTreeBuildType build_type = WaveletTreeBuildType::Standard)
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) {
228 root_ = npos;
229 return;
230 }
231 std::vector<PreWaveletNode> nodes;
232 nodes.reserve(alphabet_size_);
233 std::vector<size_t> nodes_structure;
234 nodes_structure.reserve(alphabet_size_);
235
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);
242 } else {
243 struct Node {
244 size_t size, left, right;
245 };
246 std::vector<Node> huffman_nodes(alphabet_size_, {0, 0, 0});
247 for (auto symb : data) {
248 huffman_nodes[symb].size++;
249 }
250
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);
255 }
256 while (queue.size() >= 2) {
257 auto right = queue.top().second;
258 queue.pop();
259 auto left = queue.top().second;
260 queue.pop();
261 huffman_nodes.push_back(
262 {huffman_nodes[left].size + huffman_nodes[right].size, left + 1,
263 right + 1});
264 queue.emplace(huffman_nodes.back().size, huffman_nodes.size() - 1);
265 }
266
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);
272 return 1;
273 }
274 size_t ind = nodes_structure.size(), subtree = 0;
275 if (size > 0) {
276 nodes_structure.push_back(0);
277 }
278 subtree += enumerate(left - 1);
279 if (size > 0) {
280 nodes_structure[ind] = subtree;
281 }
282 subtree += enumerate(right - 1);
283 return subtree;
284 };
285
286 permutation_.resize(alphabet_size_);
287 inverse_permutation_.reserve(alphabet_size_);
288 enumerate(huffman_nodes.size() - 1);
289 }
290
291 std::vector<size_t> prefix_sum(alphabet_size + 1);
292 for (auto symbol : data) {
293 prefix_sum[permutation_[symbol] + 1]++;
294 }
295 for (size_t i = 0; i < alphabet_size_; i++) {
296 prefix_sum[i + 1] += prefix_sum[i];
297 }
298
299 root_ = build_node(
300 0, alphabet_size_, npos,
301 [&](node_index_t node) { return nodes_structure[node]; }, prefix_sum,
302 nodes);
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;
309 if (go_right) {
310 current = node.right_child;
311 } else {
312 current = node.left_child;
313 }
314 }
315 }
316 nodes_.reserve(nodes.size());
317 for (auto& node : nodes) {
318 nodes_.emplace_back(std::move(node));
319 }
320 }
321
330 size_t rank_impl(uint64_t symbol, size_t pos) const {
331 if (symbol >= alphabet_size_) [[unlikely]] {
332 return 0;
333 }
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;
340 } else {
341 pos = node.data.rank(pos);
342 current = node.right_child;
343 }
344 }
345 return pos;
346 }
347
356 size_t select_impl(uint64_t symbol, size_t rank) const {
357 if (symbol >= alphabet_size_ || data_size_ == 0) [[unlikely]] {
358 return data_size_;
359 }
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) {
365 rank = node.data.select0(rank) + 1;
366 } else {
367 rank = node.data.select(rank) + 1;
368 }
369 }
370 return rank - 1;
371 }
372
381 std::vector<uint64_t> get_segment_impl(size_t begin, size_t end) const {
382 if (alphabet_size_ == 0 || data_size_ == 0 || begin >= end) [[unlikely]] {
383 return {};
384 }
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);
391 return result;
392 }
393
398 size_t size_impl() const { return data_size_; }
399
405 bs << alphabet_size_ << data_size_ << root_ << nodes_.size();
406 for (const WaveletNode& node : nodes_) {
407 node.serialize(bs);
408 }
409 for (const node_index_t leaf : leaves_) {
410 bs << leaf;
411 }
412 for (const size_t idx : permutation_) {
413 bs << idx;
414 }
415 }
416
417 static WaveletTreeIndex<ReadOnlyStorageView> deserialize(
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);
424 };
425 read(result.alphabet_size_);
426 read(result.data_size_);
427 read(result.root_);
428 size_t size;
429 read(size);
430 result.nodes_.resize(size);
431 for (auto& node : result.nodes_) {
432 node = WaveletNode::deserialize(data);
433 }
434 result.leaves_.resize(result.alphabet_size_);
435 for (node_index_t& leaf : result.leaves_) {
436 read(leaf);
437 }
438 result.permutation_.resize(result.alphabet_size_);
439 for (size_t& index : result.permutation_) {
440 read(index);
441 }
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;
445 }
446 return result;
447 }
448};
449
450using WaveletTree = WaveletTreeIndex<AlignedStorage>;
451using WaveletTreeView = WaveletTreeIndex<ReadOnlyStorageView>;
452
453} // namespace pixie
Owning storage rounded up to 64-byte blocks.
Definition aligned.h:37
Definition bit_stream.h:9
size_t size() const
Returns the number of written bits.
Definition bit_stream.h:57
std::uint64_t select(std::size_t rank) const
Return the position of the rank-th one bit.
Definition rank_select.h:72
std::uint64_t select0(std::size_t rank) const
Return the position of the rank-th zero bit.
Definition rank_select.h:81
std::uint64_t rank0(std::size_t end_position) const
Count zero bits in the prefix [0, end_position).
Definition rank_select.h:62
std::size_t size() const
Return the number of valid bits.
Definition rank_select.h:31
std::uint64_t rank(std::size_t end_position) const
Count one bits in the prefix [0, end_position).
Definition rank_select.h:53
Rank/select support over an external packed bit sequence.
Definition support.h:55
static ReadOnlyStorageView deserialize(std::span< const std::byte > &data)
Deserialize a size-prefixed view and advance data.
Definition read_only_view.h:44
auto writable_words64()
Return writable storage as 64-bit words.
Definition storage.h:102
CRTP facade for wavelet-tree queries.
Definition wavelet_tree.h:27
std::size_t size() const
Return the number of symbols in the indexed sequence.
Definition wavelet_tree.h:33
std::size_t rank(std::uint64_t symbol, std::size_t end_position) const
Definition wavelet_tree.h:47
Definition index.h:18
WaveletTreeIndex(size_t alphabet_size, std::span< const uint64_t > data, const WaveletTreeBuildType build_type=WaveletTreeBuildType::Standard)
Definition index.h:219
size_t select_impl(uint64_t symbol, size_t rank) const
Select the position of the rank-th specified symbol (1-indexed)
Definition index.h:356
size_t rank_impl(uint64_t symbol, size_t pos) const
Rank of specified symbol up to position pos (exclusive)
Definition index.h:330
std::vector< uint64_t > get_segment_impl(size_t begin, size_t end) const
Accumulates the original data segment.
Definition index.h:381
size_t size_impl() const
Definition index.h:398
void serialize(pixie::OutputBitStream &bs) const
Writes a wavelet tree serialization to the bit stream.
Definition index.h:404
Common interface for wavelet-tree indexes.
WaveletTreeBuildType
Construction strategy for a wavelet-tree implementation.
Definition wavelet_tree.h:18