wavelet_matrix.hpp¶
Coordinate-compressed static wavelet matrix with range order statistics.
Verified by range_kth_smallest, static_range_sum_with_upper_bound.
在静态序列上回答区间第 k 小、值域计数、前驱后继等顺序统计查询。
Implementation¶
#ifndef NOYA_WAVELET_MATRIX_HPP
#define NOYA_WAVELET_MATRIX_HPP 1
/// @complexity Time: O(n log sigma) build and O(log sigma) order/frequency query.
/// Space: O(n log sigma).
#include <algorithm>
#include <bit>
#include <cassert>
#include <optional>
#include <vector>
namespace noya {
/// @brief Coordinate-compressed static wavelet matrix with range order
/// statistics.
template <class T, class Sum = long long> struct wavelet_matrix {
int n = 0;
int levels = 0;
std::vector<T> sorted_values;
std::vector<int> middle;
std::vector<std::vector<int>> prefix_one;
std::vector<std::vector<Sum>> prefix_zero_sum;
std::vector<Sum> prefix_total;
wavelet_matrix() = default;
explicit wavelet_matrix(const std::vector<T> &values) { build(values); }
/// @brief Rebuild the static matrix from values.
void build(const std::vector<T> &values) {
n = int(values.size());
sorted_values = values;
std::sort(sorted_values.begin(), sorted_values.end());
sorted_values.erase(std::unique(sorted_values.begin(), sorted_values.end()),
sorted_values.end());
int alphabet_size = int(sorted_values.size());
levels = std::max(
1, int(std::bit_width(unsigned(std::max(0, alphabet_size - 1)))));
middle.assign(levels, 0);
prefix_one.assign(levels, std::vector<int>(n + 1));
prefix_zero_sum.assign(levels, std::vector<Sum>(n + 1));
prefix_total.assign(n + 1, Sum{});
for (int i = 0; i < n; i++) {
prefix_total[i + 1] = prefix_total[i] + Sum(values[i]);
}
std::vector<int> rank(n);
for (int i = 0; i < n; i++) {
rank[i] = int(std::lower_bound(sorted_values.begin(), sorted_values.end(),
values[i]) -
sorted_values.begin());
}
std::vector<T> current_values = values;
for (int level = 0; level < levels; level++) {
int bit = levels - 1 - level;
std::vector<int> zero_rank, one_rank;
std::vector<T> zero_value, one_value;
zero_rank.reserve(n);
one_rank.reserve(n);
zero_value.reserve(n);
one_value.reserve(n);
for (int i = 0; i < n; i++) {
bool one = (rank[i] >> bit) & 1;
prefix_one[level][i + 1] = prefix_one[level][i] + one;
prefix_zero_sum[level][i + 1] =
prefix_zero_sum[level][i] + (one ? Sum{} : Sum(current_values[i]));
if (one) {
one_rank.push_back(rank[i]);
one_value.push_back(current_values[i]);
} else {
zero_rank.push_back(rank[i]);
zero_value.push_back(current_values[i]);
}
}
middle[level] = int(zero_rank.size());
zero_rank.insert(zero_rank.end(), one_rank.begin(), one_rank.end());
zero_value.insert(zero_value.end(), one_value.begin(), one_value.end());
rank.swap(zero_rank);
current_values.swap(zero_value);
}
}
/// @brief Return the number of stored values.
int size() const { return n; }
/// @brief Return the k-th smallest value in values[l, r), where k is
/// zero-indexed.
T kth_smallest(int l, int r, int k) const {
assert(0 <= l && l <= r && r <= n);
assert(0 <= k && k < r - l);
unsigned rank = 0;
for (int level = 0; level < levels; level++) {
int ones_l = prefix_one[level][l];
int ones_r = prefix_one[level][r];
int zeros_l = l - ones_l;
int zeros_r = r - ones_r;
int zero_count = zeros_r - zeros_l;
if (k < zero_count) {
l = zeros_l;
r = zeros_r;
} else {
k -= zero_count;
rank |= 1U << (levels - 1 - level);
l = middle[level] + ones_l;
r = middle[level] + ones_r;
}
}
assert(rank < sorted_values.size());
return sorted_values[rank];
}
/// @brief Count values x in values[l, r) satisfying x < upper.
int count_less(int l, int r, const T &upper) const {
check_range(l, r);
int rank = int(
std::lower_bound(sorted_values.begin(), sorted_values.end(), upper) -
sorted_values.begin());
return count_less_rank(l, r, rank);
}
/// @brief Count values x in values[l, r) satisfying lower <= x < upper.
int range_freq(int l, int r, const T &lower, const T &upper) const {
assert(!(upper < lower));
return count_less(l, r, upper) - count_less(l, r, lower);
}
/// @brief Count occurrences of value in values[l, r).
int count(int l, int r, const T &value) const {
check_range(l, r);
int lower = int(
std::lower_bound(sorted_values.begin(), sorted_values.end(), value) -
sorted_values.begin());
int upper = int(
std::upper_bound(sorted_values.begin(), sorted_values.end(), value) -
sorted_values.begin());
return count_less_rank(l, r, upper) - count_less_rank(l, r, lower);
}
/// @brief Sum values x in values[l, r) satisfying lower <= x < upper.
Sum range_sum(int l, int r, const T &lower, const T &upper) const {
assert(!(upper < lower));
return sum_less(l, r, upper) - sum_less(l, r, lower);
}
/// @brief Return the largest value below upper in values[l, r), or nullopt.
std::optional<T> prev_value(int l, int r, const T &upper) const {
int count = count_less(l, r, upper);
if (count == 0) {
return std::nullopt;
}
return kth_smallest(l, r, count - 1);
}
/// @brief Return the smallest value at least lower in values[l, r), or
/// nullopt.
std::optional<T> next_value(int l, int r, const T &lower) const {
int count = count_less(l, r, lower);
if (count == r - l) {
return std::nullopt;
}
return kth_smallest(l, r, count);
}
private:
void check_range(int l, int r) const { assert(0 <= l && l <= r && r <= n); }
int count_less_rank(int l, int r, int rank) const {
check_range(l, r);
if (rank <= 0) {
return 0;
}
if (rank >= int(sorted_values.size())) {
return r - l;
}
int result = 0;
for (int level = 0; level < levels; level++) {
int ones_l = prefix_one[level][l];
int ones_r = prefix_one[level][r];
int zeros_l = l - ones_l;
int zeros_r = r - ones_r;
int bit = levels - 1 - level;
if ((rank >> bit) & 1) {
result += zeros_r - zeros_l;
l = middle[level] + ones_l;
r = middle[level] + ones_r;
} else {
l = zeros_l;
r = zeros_r;
}
}
return result;
}
Sum sum_less(int l, int r, const T &upper) const {
check_range(l, r);
int rank = int(
std::lower_bound(sorted_values.begin(), sorted_values.end(), upper) -
sorted_values.begin());
if (rank <= 0) {
return Sum{};
}
if (rank >= int(sorted_values.size())) {
return prefix_total[r] - prefix_total[l];
}
Sum result{};
for (int level = 0; level < levels; level++) {
int ones_l = prefix_one[level][l];
int ones_r = prefix_one[level][r];
int zeros_l = l - ones_l;
int zeros_r = r - ones_r;
int bit = levels - 1 - level;
if ((rank >> bit) & 1) {
result += prefix_zero_sum[level][r] - prefix_zero_sum[level][l];
l = middle[level] + ones_l;
r = middle[level] + ones_r;
} else {
l = zeros_l;
r = zeros_r;
}
}
return result;
}
};
} // namespace noya
#endif // NOYA_WAVELET_MATRIX_HPP
#include <algorithm>
#include <bit>
#include <cassert>
#include <optional>
#include <vector>
/// @complexity Time: O(n log sigma) build and O(log sigma) order/frequency query.
/// Space: O(n log sigma).
namespace noya {
/// @brief Coordinate-compressed static wavelet matrix with range order
/// statistics.
template <class T, class Sum = long long> struct wavelet_matrix {
int n = 0;
int levels = 0;
std::vector<T> sorted_values;
std::vector<int> middle;
std::vector<std::vector<int>> prefix_one;
std::vector<std::vector<Sum>> prefix_zero_sum;
std::vector<Sum> prefix_total;
wavelet_matrix() = default;
explicit wavelet_matrix(const std::vector<T> &values) { build(values); }
/// @brief Rebuild the static matrix from values.
void build(const std::vector<T> &values) {
n = int(values.size());
sorted_values = values;
std::sort(sorted_values.begin(), sorted_values.end());
sorted_values.erase(std::unique(sorted_values.begin(), sorted_values.end()),
sorted_values.end());
int alphabet_size = int(sorted_values.size());
levels = std::max(
1, int(std::bit_width(unsigned(std::max(0, alphabet_size - 1)))));
middle.assign(levels, 0);
prefix_one.assign(levels, std::vector<int>(n + 1));
prefix_zero_sum.assign(levels, std::vector<Sum>(n + 1));
prefix_total.assign(n + 1, Sum{});
for (int i = 0; i < n; i++) {
prefix_total[i + 1] = prefix_total[i] + Sum(values[i]);
}
std::vector<int> rank(n);
for (int i = 0; i < n; i++) {
rank[i] = int(std::lower_bound(sorted_values.begin(), sorted_values.end(),
values[i]) -
sorted_values.begin());
}
std::vector<T> current_values = values;
for (int level = 0; level < levels; level++) {
int bit = levels - 1 - level;
std::vector<int> zero_rank, one_rank;
std::vector<T> zero_value, one_value;
zero_rank.reserve(n);
one_rank.reserve(n);
zero_value.reserve(n);
one_value.reserve(n);
for (int i = 0; i < n; i++) {
bool one = (rank[i] >> bit) & 1;
prefix_one[level][i + 1] = prefix_one[level][i] + one;
prefix_zero_sum[level][i + 1] =
prefix_zero_sum[level][i] + (one ? Sum{} : Sum(current_values[i]));
if (one) {
one_rank.push_back(rank[i]);
one_value.push_back(current_values[i]);
} else {
zero_rank.push_back(rank[i]);
zero_value.push_back(current_values[i]);
}
}
middle[level] = int(zero_rank.size());
zero_rank.insert(zero_rank.end(), one_rank.begin(), one_rank.end());
zero_value.insert(zero_value.end(), one_value.begin(), one_value.end());
rank.swap(zero_rank);
current_values.swap(zero_value);
}
}
/// @brief Return the number of stored values.
int size() const { return n; }
/// @brief Return the k-th smallest value in values[l, r), where k is
/// zero-indexed.
T kth_smallest(int l, int r, int k) const {
assert(0 <= l && l <= r && r <= n);
assert(0 <= k && k < r - l);
unsigned rank = 0;
for (int level = 0; level < levels; level++) {
int ones_l = prefix_one[level][l];
int ones_r = prefix_one[level][r];
int zeros_l = l - ones_l;
int zeros_r = r - ones_r;
int zero_count = zeros_r - zeros_l;
if (k < zero_count) {
l = zeros_l;
r = zeros_r;
} else {
k -= zero_count;
rank |= 1U << (levels - 1 - level);
l = middle[level] + ones_l;
r = middle[level] + ones_r;
}
}
assert(rank < sorted_values.size());
return sorted_values[rank];
}
/// @brief Count values x in values[l, r) satisfying x < upper.
int count_less(int l, int r, const T &upper) const {
check_range(l, r);
int rank = int(
std::lower_bound(sorted_values.begin(), sorted_values.end(), upper) -
sorted_values.begin());
return count_less_rank(l, r, rank);
}
/// @brief Count values x in values[l, r) satisfying lower <= x < upper.
int range_freq(int l, int r, const T &lower, const T &upper) const {
assert(!(upper < lower));
return count_less(l, r, upper) - count_less(l, r, lower);
}
/// @brief Count occurrences of value in values[l, r).
int count(int l, int r, const T &value) const {
check_range(l, r);
int lower = int(
std::lower_bound(sorted_values.begin(), sorted_values.end(), value) -
sorted_values.begin());
int upper = int(
std::upper_bound(sorted_values.begin(), sorted_values.end(), value) -
sorted_values.begin());
return count_less_rank(l, r, upper) - count_less_rank(l, r, lower);
}
/// @brief Sum values x in values[l, r) satisfying lower <= x < upper.
Sum range_sum(int l, int r, const T &lower, const T &upper) const {
assert(!(upper < lower));
return sum_less(l, r, upper) - sum_less(l, r, lower);
}
/// @brief Return the largest value below upper in values[l, r), or nullopt.
std::optional<T> prev_value(int l, int r, const T &upper) const {
int count = count_less(l, r, upper);
if (count == 0) {
return std::nullopt;
}
return kth_smallest(l, r, count - 1);
}
/// @brief Return the smallest value at least lower in values[l, r), or
/// nullopt.
std::optional<T> next_value(int l, int r, const T &lower) const {
int count = count_less(l, r, lower);
if (count == r - l) {
return std::nullopt;
}
return kth_smallest(l, r, count);
}
private:
void check_range(int l, int r) const { assert(0 <= l && l <= r && r <= n); }
int count_less_rank(int l, int r, int rank) const {
check_range(l, r);
if (rank <= 0) {
return 0;
}
if (rank >= int(sorted_values.size())) {
return r - l;
}
int result = 0;
for (int level = 0; level < levels; level++) {
int ones_l = prefix_one[level][l];
int ones_r = prefix_one[level][r];
int zeros_l = l - ones_l;
int zeros_r = r - ones_r;
int bit = levels - 1 - level;
if ((rank >> bit) & 1) {
result += zeros_r - zeros_l;
l = middle[level] + ones_l;
r = middle[level] + ones_r;
} else {
l = zeros_l;
r = zeros_r;
}
}
return result;
}
Sum sum_less(int l, int r, const T &upper) const {
check_range(l, r);
int rank = int(
std::lower_bound(sorted_values.begin(), sorted_values.end(), upper) -
sorted_values.begin());
if (rank <= 0) {
return Sum{};
}
if (rank >= int(sorted_values.size())) {
return prefix_total[r] - prefix_total[l];
}
Sum result{};
for (int level = 0; level < levels; level++) {
int ones_l = prefix_one[level][l];
int ones_r = prefix_one[level][r];
int zeros_l = l - ones_l;
int zeros_r = r - ones_r;
int bit = levels - 1 - level;
if ((rank >> bit) & 1) {
result += prefix_zero_sum[level][r] - prefix_zero_sum[level][l];
l = middle[level] + ones_l;
r = middle[level] + ones_r;
} else {
l = zeros_l;
r = zeros_r;
}
}
return result;
}
};
} // namespace noya