Skip to content

wavelet_matrix.hpp

SECTIONData Structure INCLUDEnoya/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

跳到代码 · GitHub ↗

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