wavelet_matrix.hpp¶
在静态序列上回答区间第 \(k\) 小、值域计数、前驱后继等顺序统计查询。
Complexity: Time: O(n log sigma) build and O(log sigma) order/frequency query. Space: O(n log sigma).
AC 记录:range_kth_smallest, static_range_sum_with_upper_bound。
Implementation¶
当前头文件,省略 include guard;依赖见 #include。
/// @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 lg = 0;
std::vector<T> xs;
std::vector<int> mid;
std::vector<std::vector<int>> s1;
std::vector<std::vector<Sum>> s0;
std::vector<Sum> s;
wavelet_matrix() = default;
explicit wavelet_matrix(const std::vector<T> &a) { build(a); }
/// @brief Rebuild the static matrix from values.
void build(const std::vector<T> &a) {
n = int(a.size());
xs = a;
std::sort(xs.begin(), xs.end());
xs.erase(std::unique(xs.begin(), xs.end()), xs.end());
int m = int(xs.size());
lg = std::max(1, int(std::bit_width(unsigned(std::max(0, m - 1)))));
mid.assign(lg, 0);
s1.assign(lg, std::vector<int>(n + 1));
s0.assign(lg, std::vector<Sum>(n + 1));
s.assign(n + 1, Sum{});
for (int i = 0; i < n; i++) {
s[i + 1] = s[i] + Sum(a[i]);
}
std::vector<int> rk(n);
for (int i = 0; i < n; i++) {
rk[i] = int(std::lower_bound(xs.begin(), xs.end(), a[i]) - xs.begin());
}
std::vector<T> a0 = a;
for (int dep = 0; dep < lg; dep++) {
int bit = lg - 1 - dep;
std::vector<int> r0, r1;
std::vector<T> v0, v1;
r0.reserve(n);
r1.reserve(n);
v0.reserve(n);
v1.reserve(n);
for (int i = 0; i < n; i++) {
bool one = (rk[i] >> bit) & 1;
s1[dep][i + 1] = s1[dep][i] + one;
s0[dep][i + 1] = s0[dep][i] + (one ? Sum{} : Sum(a0[i]));
if (one) {
r1.push_back(rk[i]);
v1.push_back(a0[i]);
} else {
r0.push_back(rk[i]);
v0.push_back(a0[i]);
}
}
mid[dep] = int(r0.size());
r0.insert(r0.end(), r1.begin(), r1.end());
v0.insert(v0.end(), v1.begin(), v1.end());
rk.swap(r0);
a0.swap(v0);
}
}
/// @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 rk = 0;
for (int dep = 0; dep < lg; dep++) {
int ol = s1[dep][l];
int or_ = s1[dep][r];
int zl = l - ol;
int zr = r - or_;
int c0 = zr - zl;
if (k < c0) {
l = zl;
r = zr;
} else {
k -= c0;
rk |= 1U << (lg - 1 - dep);
l = mid[dep] + ol;
r = mid[dep] + or_;
}
}
assert(rk < xs.size());
return xs[rk];
}
/// @brief Count values x at positions [l, r) satisfying x < hi.
int count_less(int l, int r, const T &hi) const {
check_range(l, r);
int rk = int(std::lower_bound(xs.begin(), xs.end(), hi) - xs.begin());
return count_less_rank(l, r, rk);
}
/// @brief Count values x at positions [l, r) satisfying lo <= x < hi.
int range_freq(int l, int r, const T &lo, const T &hi) const {
assert(!(hi < lo));
return count_less(l, r, hi) - count_less(l, r, lo);
}
/// @brief Count occurrences of val at positions [l, r).
int count(int l, int r, const T &val) const {
check_range(l, r);
int lo = int(std::lower_bound(xs.begin(), xs.end(), val) - xs.begin());
int hi = int(std::upper_bound(xs.begin(), xs.end(), val) - xs.begin());
return count_less_rank(l, r, hi) - count_less_rank(l, r, lo);
}
/// @brief Sum values x at positions [l, r) satisfying lo <= x < hi.
Sum range_sum(int l, int r, const T &lo, const T &hi) const {
assert(!(hi < lo));
return sum_less(l, r, hi) - sum_less(l, r, lo);
}
/// @brief Return the largest value below hi at positions [l, r), or nullopt.
std::optional<T> prev_value(int l, int r, const T &hi) const {
int cnt = count_less(l, r, hi);
if (cnt == 0) {
return std::nullopt;
}
return kth_smallest(l, r, cnt - 1);
}
/// @brief Return the smallest value at least lo at positions [l, r), or
/// nullopt.
std::optional<T> next_value(int l, int r, const T &lo) const {
int cnt = count_less(l, r, lo);
if (cnt == r - l) {
return std::nullopt;
}
return kth_smallest(l, r, cnt);
}
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 rk) const {
check_range(l, r);
if (rk <= 0) {
return 0;
}
if (rk >= int(xs.size())) {
return r - l;
}
int res = 0;
for (int dep = 0; dep < lg; dep++) {
int ol = s1[dep][l];
int or_ = s1[dep][r];
int zl = l - ol;
int zr = r - or_;
int bit = lg - 1 - dep;
if ((rk >> bit) & 1) {
res += zr - zl;
l = mid[dep] + ol;
r = mid[dep] + or_;
} else {
l = zl;
r = zr;
}
}
return res;
}
Sum sum_less(int l, int r, const T &hi) const {
check_range(l, r);
int rk = int(std::lower_bound(xs.begin(), xs.end(), hi) - xs.begin());
if (rk <= 0) {
return Sum{};
}
if (rk >= int(xs.size())) {
return s[r] - s[l];
}
Sum res{};
for (int dep = 0; dep < lg; dep++) {
int ol = s1[dep][l];
int or_ = s1[dep][r];
int zl = l - ol;
int zr = r - or_;
int bit = lg - 1 - dep;
if ((rk >> bit) & 1) {
res += s0[dep][r] - s0[dep][l];
l = mid[dep] + ol;
r = mid[dep] + or_;
} else {
l = zl;
r = zr;
}
}
return res;
}
};
} // namespace noya
#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 lg = 0;
std::vector<T> xs;
std::vector<int> mid;
std::vector<std::vector<int>> s1;
std::vector<std::vector<Sum>> s0;
std::vector<Sum> s;
wavelet_matrix() = default;
explicit wavelet_matrix(const std::vector<T> &a) { build(a); }
/// @brief Rebuild the static matrix from values.
void build(const std::vector<T> &a) {
n = int(a.size());
xs = a;
std::sort(xs.begin(), xs.end());
xs.erase(std::unique(xs.begin(), xs.end()), xs.end());
int m = int(xs.size());
lg = std::max(1, int(std::bit_width(unsigned(std::max(0, m - 1)))));
mid.assign(lg, 0);
s1.assign(lg, std::vector<int>(n + 1));
s0.assign(lg, std::vector<Sum>(n + 1));
s.assign(n + 1, Sum{});
for (int i = 0; i < n; i++) {
s[i + 1] = s[i] + Sum(a[i]);
}
std::vector<int> rk(n);
for (int i = 0; i < n; i++) {
rk[i] = int(std::lower_bound(xs.begin(), xs.end(), a[i]) - xs.begin());
}
std::vector<T> a0 = a;
for (int dep = 0; dep < lg; dep++) {
int bit = lg - 1 - dep;
std::vector<int> r0, r1;
std::vector<T> v0, v1;
r0.reserve(n);
r1.reserve(n);
v0.reserve(n);
v1.reserve(n);
for (int i = 0; i < n; i++) {
bool one = (rk[i] >> bit) & 1;
s1[dep][i + 1] = s1[dep][i] + one;
s0[dep][i + 1] = s0[dep][i] + (one ? Sum{} : Sum(a0[i]));
if (one) {
r1.push_back(rk[i]);
v1.push_back(a0[i]);
} else {
r0.push_back(rk[i]);
v0.push_back(a0[i]);
}
}
mid[dep] = int(r0.size());
r0.insert(r0.end(), r1.begin(), r1.end());
v0.insert(v0.end(), v1.begin(), v1.end());
rk.swap(r0);
a0.swap(v0);
}
}
/// @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 rk = 0;
for (int dep = 0; dep < lg; dep++) {
int ol = s1[dep][l];
int or_ = s1[dep][r];
int zl = l - ol;
int zr = r - or_;
int c0 = zr - zl;
if (k < c0) {
l = zl;
r = zr;
} else {
k -= c0;
rk |= 1U << (lg - 1 - dep);
l = mid[dep] + ol;
r = mid[dep] + or_;
}
}
assert(rk < xs.size());
return xs[rk];
}
/// @brief Count values x at positions [l, r) satisfying x < hi.
int count_less(int l, int r, const T &hi) const {
check_range(l, r);
int rk = int(std::lower_bound(xs.begin(), xs.end(), hi) - xs.begin());
return count_less_rank(l, r, rk);
}
/// @brief Count values x at positions [l, r) satisfying lo <= x < hi.
int range_freq(int l, int r, const T &lo, const T &hi) const {
assert(!(hi < lo));
return count_less(l, r, hi) - count_less(l, r, lo);
}
/// @brief Count occurrences of val at positions [l, r).
int count(int l, int r, const T &val) const {
check_range(l, r);
int lo = int(std::lower_bound(xs.begin(), xs.end(), val) - xs.begin());
int hi = int(std::upper_bound(xs.begin(), xs.end(), val) - xs.begin());
return count_less_rank(l, r, hi) - count_less_rank(l, r, lo);
}
/// @brief Sum values x at positions [l, r) satisfying lo <= x < hi.
Sum range_sum(int l, int r, const T &lo, const T &hi) const {
assert(!(hi < lo));
return sum_less(l, r, hi) - sum_less(l, r, lo);
}
/// @brief Return the largest value below hi at positions [l, r), or nullopt.
std::optional<T> prev_value(int l, int r, const T &hi) const {
int cnt = count_less(l, r, hi);
if (cnt == 0) {
return std::nullopt;
}
return kth_smallest(l, r, cnt - 1);
}
/// @brief Return the smallest value at least lo at positions [l, r), or
/// nullopt.
std::optional<T> next_value(int l, int r, const T &lo) const {
int cnt = count_less(l, r, lo);
if (cnt == r - l) {
return std::nullopt;
}
return kth_smallest(l, r, cnt);
}
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 rk) const {
check_range(l, r);
if (rk <= 0) {
return 0;
}
if (rk >= int(xs.size())) {
return r - l;
}
int res = 0;
for (int dep = 0; dep < lg; dep++) {
int ol = s1[dep][l];
int or_ = s1[dep][r];
int zl = l - ol;
int zr = r - or_;
int bit = lg - 1 - dep;
if ((rk >> bit) & 1) {
res += zr - zl;
l = mid[dep] + ol;
r = mid[dep] + or_;
} else {
l = zl;
r = zr;
}
}
return res;
}
Sum sum_less(int l, int r, const T &hi) const {
check_range(l, r);
int rk = int(std::lower_bound(xs.begin(), xs.end(), hi) - xs.begin());
if (rk <= 0) {
return Sum{};
}
if (rk >= int(xs.size())) {
return s[r] - s[l];
}
Sum res{};
for (int dep = 0; dep < lg; dep++) {
int ol = s1[dep][l];
int or_ = s1[dep][r];
int zl = l - ol;
int zr = r - or_;
int bit = lg - 1 - dep;
if ((rk >> bit) & 1) {
res += s0[dep][r] - s0[dep][l];
l = mid[dep] + ol;
r = mid[dep] + or_;
} else {
l = zl;
r = zr;
}
}
return res;
}
};
} // 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 lg = 0;
std::vector<T> xs;
std::vector<int> mid;
std::vector<std::vector<int>> s1;
std::vector<std::vector<Sum>> s0;
std::vector<Sum> s;
wavelet_matrix() = default;
explicit wavelet_matrix(const std::vector<T> &a) { build(a); }
/// @brief Rebuild the static matrix from values.
void build(const std::vector<T> &a) {
n = int(a.size());
xs = a;
std::sort(xs.begin(), xs.end());
xs.erase(std::unique(xs.begin(), xs.end()), xs.end());
int m = int(xs.size());
lg = std::max(1, int(std::bit_width(unsigned(std::max(0, m - 1)))));
mid.assign(lg, 0);
s1.assign(lg, std::vector<int>(n + 1));
s0.assign(lg, std::vector<Sum>(n + 1));
s.assign(n + 1, Sum{});
for (int i = 0; i < n; i++) {
s[i + 1] = s[i] + Sum(a[i]);
}
std::vector<int> rk(n);
for (int i = 0; i < n; i++) {
rk[i] = int(std::lower_bound(xs.begin(), xs.end(), a[i]) - xs.begin());
}
std::vector<T> a0 = a;
for (int dep = 0; dep < lg; dep++) {
int bit = lg - 1 - dep;
std::vector<int> r0, r1;
std::vector<T> v0, v1;
r0.reserve(n);
r1.reserve(n);
v0.reserve(n);
v1.reserve(n);
for (int i = 0; i < n; i++) {
bool one = (rk[i] >> bit) & 1;
s1[dep][i + 1] = s1[dep][i] + one;
s0[dep][i + 1] = s0[dep][i] + (one ? Sum{} : Sum(a0[i]));
if (one) {
r1.push_back(rk[i]);
v1.push_back(a0[i]);
} else {
r0.push_back(rk[i]);
v0.push_back(a0[i]);
}
}
mid[dep] = int(r0.size());
r0.insert(r0.end(), r1.begin(), r1.end());
v0.insert(v0.end(), v1.begin(), v1.end());
rk.swap(r0);
a0.swap(v0);
}
}
/// @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 rk = 0;
for (int dep = 0; dep < lg; dep++) {
int ol = s1[dep][l];
int or_ = s1[dep][r];
int zl = l - ol;
int zr = r - or_;
int c0 = zr - zl;
if (k < c0) {
l = zl;
r = zr;
} else {
k -= c0;
rk |= 1U << (lg - 1 - dep);
l = mid[dep] + ol;
r = mid[dep] + or_;
}
}
assert(rk < xs.size());
return xs[rk];
}
/// @brief Count values x at positions [l, r) satisfying x < hi.
int count_less(int l, int r, const T &hi) const {
check_range(l, r);
int rk = int(std::lower_bound(xs.begin(), xs.end(), hi) - xs.begin());
return count_less_rank(l, r, rk);
}
/// @brief Count values x at positions [l, r) satisfying lo <= x < hi.
int range_freq(int l, int r, const T &lo, const T &hi) const {
assert(!(hi < lo));
return count_less(l, r, hi) - count_less(l, r, lo);
}
/// @brief Count occurrences of val at positions [l, r).
int count(int l, int r, const T &val) const {
check_range(l, r);
int lo = int(std::lower_bound(xs.begin(), xs.end(), val) - xs.begin());
int hi = int(std::upper_bound(xs.begin(), xs.end(), val) - xs.begin());
return count_less_rank(l, r, hi) - count_less_rank(l, r, lo);
}
/// @brief Sum values x at positions [l, r) satisfying lo <= x < hi.
Sum range_sum(int l, int r, const T &lo, const T &hi) const {
assert(!(hi < lo));
return sum_less(l, r, hi) - sum_less(l, r, lo);
}
/// @brief Return the largest value below hi at positions [l, r), or nullopt.
std::optional<T> prev_value(int l, int r, const T &hi) const {
int cnt = count_less(l, r, hi);
if (cnt == 0) {
return std::nullopt;
}
return kth_smallest(l, r, cnt - 1);
}
/// @brief Return the smallest value at least lo at positions [l, r), or
/// nullopt.
std::optional<T> next_value(int l, int r, const T &lo) const {
int cnt = count_less(l, r, lo);
if (cnt == r - l) {
return std::nullopt;
}
return kth_smallest(l, r, cnt);
}
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 rk) const {
check_range(l, r);
if (rk <= 0) {
return 0;
}
if (rk >= int(xs.size())) {
return r - l;
}
int res = 0;
for (int dep = 0; dep < lg; dep++) {
int ol = s1[dep][l];
int or_ = s1[dep][r];
int zl = l - ol;
int zr = r - or_;
int bit = lg - 1 - dep;
if ((rk >> bit) & 1) {
res += zr - zl;
l = mid[dep] + ol;
r = mid[dep] + or_;
} else {
l = zl;
r = zr;
}
}
return res;
}
Sum sum_less(int l, int r, const T &hi) const {
check_range(l, r);
int rk = int(std::lower_bound(xs.begin(), xs.end(), hi) - xs.begin());
if (rk <= 0) {
return Sum{};
}
if (rk >= int(xs.size())) {
return s[r] - s[l];
}
Sum res{};
for (int dep = 0; dep < lg; dep++) {
int ol = s1[dep][l];
int or_ = s1[dep][r];
int zl = l - ol;
int zr = r - or_;
int bit = lg - 1 - dep;
if ((rk >> bit) & 1) {
res += s0[dep][r] - s0[dep][l];
l = mid[dep] + ol;
r = mid[dep] + or_;
} else {
l = zl;
r = zr;
}
}
return res;
}
};
} // namespace noya