Skip to content

wavelet_matrix.hpp

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

View on GitHub

#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