Skip to content

segment_tree_beats.hpp

SECTIONData Structure INCLUDEnoya/segment_tree_beats.hpp

Segment Tree Beats supporting range add, chmin, chmax, sum, minimum, and maximum queries; bound updates run in amortized O(log^2 n) time.

Verified by range_add_range_min, range_chmin_chmax_add_range_sum.

维护区间 chmin、chmax、加法与区间和/最值;适合普通懒标记无法直接表示的截断更新。

Implementation

View on GitHub

#ifndef NOYA_SEGMENT_TREE_BEATS_HPP
#define NOYA_SEGMENT_TREE_BEATS_HPP 1

/// @complexity Time: O(n) build; amortized O(log^2 n) chmin/chmax and O(log n) add/query.
/// Space: O(n).

#include <algorithm>
#include <cassert>
#include <limits>
#include <vector>

namespace noya {

/// @brief Segment Tree Beats supporting range add, chmin, chmax, sum, minimum,
/// and maximum queries; bound updates run in amortized O(log^2 n) time.
template <class T> struct segment_tree_beats {
  struct node {
    T sum{};
    T maximum{};
    T second_maximum{};
    T minimum{};
    T second_minimum{};
    T lazy_add{};
    int maximum_count = 0;
    int minimum_count = 0;
    int length = 0;
    bool has_second_maximum = false;
    bool has_second_minimum = false;
  };

  int n = 0;
  std::vector<node> tree;

  segment_tree_beats() = default;
  explicit segment_tree_beats(const std::vector<T> &values) { build(values); }

  /// @brief Rebuild from values in O(n) time.
  void build(const std::vector<T> &values) {
    n = int(values.size());
    tree.assign(std::max(1, 4 * n), {});
    if (n > 0) {
      build_at(1, 0, n, values);
    }
  }

  /// @brief Add value to every element of [left, right).
  void range_add(int left, int right, const T &value) {
    check_range(left, right);
    if (left < right) {
      range_add_at(1, 0, n, left, right, value);
    }
  }

  /// @brief Replace each x in [left, right) by min(x, upper).
  void range_chmin(int left, int right, const T &upper) {
    check_range(left, right);
    if (left < right) {
      range_chmin_at(1, 0, n, left, right, upper);
    }
  }

  /// @brief Replace each x in [left, right) by max(x, lower).
  void range_chmax(int left, int right, const T &lower) {
    check_range(left, right);
    if (left < right) {
      range_chmax_at(1, 0, n, left, right, lower);
    }
  }

  /// @brief Return the sum over [left, right); an empty range has sum zero.
  T range_sum(int left, int right) {
    check_range(left, right);
    return left == right ? T{} : range_sum_at(1, 0, n, left, right);
  }

  /// @brief Return the minimum over the nonempty range [left, right).
  T range_min(int left, int right) {
    check_nonempty_range(left, right);
    return range_min_at(1, 0, n, left, right);
  }

  /// @brief Return the maximum over the nonempty range [left, right).
  T range_max(int left, int right) {
    check_nonempty_range(left, right);
    return range_max_at(1, 0, n, left, right);
  }

  /// @brief Return the sum of the entire array.
  T all_sum() const { return n == 0 ? T{} : tree[1].sum; }

  /// @brief Return the minimum of the nonempty array.
  T all_min() const {
    assert(n > 0);
    return tree[1].minimum;
  }

  /// @brief Return the maximum of the nonempty array.
  T all_max() const {
    assert(n > 0);
    return tree[1].maximum;
  }

private:
  void check_range(int left, int right) const {
    assert(0 <= left && left <= right && right <= n);
  }

  void check_nonempty_range(int left, int right) const {
    check_range(left, right);
    assert(left < right);
  }

  static node leaf(const T &value) {
    node result;
    result.sum = result.maximum = result.minimum = value;
    result.maximum_count = result.minimum_count = result.length = 1;
    return result;
  }

  static void assign_maximum(node &result, const node &left,
                             const node &right) {
    if (left.maximum == right.maximum) {
      result.maximum = left.maximum;
      result.maximum_count = left.maximum_count + right.maximum_count;
      if (left.has_second_maximum && right.has_second_maximum) {
        result.second_maximum =
            std::max(left.second_maximum, right.second_maximum);
        result.has_second_maximum = true;
      } else if (left.has_second_maximum) {
        result.second_maximum = left.second_maximum;
        result.has_second_maximum = true;
      } else if (right.has_second_maximum) {
        result.second_maximum = right.second_maximum;
        result.has_second_maximum = true;
      }
      return;
    }
    const node &higher = left.maximum > right.maximum ? left : right;
    const node &lower = left.maximum > right.maximum ? right : left;
    result.maximum = higher.maximum;
    result.maximum_count = higher.maximum_count;
    result.second_maximum =
        higher.has_second_maximum
            ? std::max(higher.second_maximum, lower.maximum)
            : lower.maximum;
    result.has_second_maximum = true;
  }

  static void assign_minimum(node &result, const node &left,
                             const node &right) {
    if (left.minimum == right.minimum) {
      result.minimum = left.minimum;
      result.minimum_count = left.minimum_count + right.minimum_count;
      if (left.has_second_minimum && right.has_second_minimum) {
        result.second_minimum =
            std::min(left.second_minimum, right.second_minimum);
        result.has_second_minimum = true;
      } else if (left.has_second_minimum) {
        result.second_minimum = left.second_minimum;
        result.has_second_minimum = true;
      } else if (right.has_second_minimum) {
        result.second_minimum = right.second_minimum;
        result.has_second_minimum = true;
      }
      return;
    }
    const node &lower = left.minimum < right.minimum ? left : right;
    const node &higher = left.minimum < right.minimum ? right : left;
    result.minimum = lower.minimum;
    result.minimum_count = lower.minimum_count;
    result.second_minimum =
        lower.has_second_minimum
            ? std::min(lower.second_minimum, higher.minimum)
            : higher.minimum;
    result.has_second_minimum = true;
  }

  static node merge_nodes(const node &left, const node &right) {
    node result;
    result.sum = left.sum + right.sum;
    result.length = left.length + right.length;
    assign_maximum(result, left, right);
    assign_minimum(result, left, right);
    return result;
  }

  void build_at(int id, int left, int right, const std::vector<T> &values) {
    if (right - left == 1) {
      tree[id] = leaf(values[left]);
      return;
    }
    int middle = (left + right) / 2;
    build_at(id * 2, left, middle, values);
    build_at(id * 2 + 1, middle, right, values);
    pull(id);
  }

  void pull(int id) { tree[id] = merge_nodes(tree[id * 2], tree[id * 2 + 1]); }

  void apply_add(int id, const T &value) {
    node &current = tree[id];
    current.sum += value * current.length;
    current.maximum += value;
    current.minimum += value;
    if (current.has_second_maximum) {
      current.second_maximum += value;
    }
    if (current.has_second_minimum) {
      current.second_minimum += value;
    }
    current.lazy_add += value;
  }

  void apply_chmin(int id, const T &upper) {
    node &current = tree[id];
    if (current.maximum <= upper) {
      return;
    }
    T old_maximum = current.maximum;
    current.sum += (upper - old_maximum) * current.maximum_count;
    if (current.minimum == old_maximum) {
      current.minimum = upper;
    } else if (current.has_second_minimum &&
               current.second_minimum == old_maximum) {
      current.second_minimum = upper;
    }
    current.maximum = upper;
  }

  void apply_chmax(int id, const T &lower) {
    node &current = tree[id];
    if (lower <= current.minimum) {
      return;
    }
    T old_minimum = current.minimum;
    current.sum += (lower - old_minimum) * current.minimum_count;
    if (current.maximum == old_minimum) {
      current.maximum = lower;
    } else if (current.has_second_maximum &&
               current.second_maximum == old_minimum) {
      current.second_maximum = lower;
    }
    current.minimum = lower;
  }

  void push(int id) {
    node &current = tree[id];
    if (current.length == 1) {
      current.lazy_add = T{};
      return;
    }
    if (current.lazy_add != T{}) {
      apply_add(id * 2, current.lazy_add);
      apply_add(id * 2 + 1, current.lazy_add);
      current.lazy_add = T{};
    }
    if (tree[id * 2].maximum > current.maximum) {
      apply_chmin(id * 2, current.maximum);
    }
    if (tree[id * 2 + 1].maximum > current.maximum) {
      apply_chmin(id * 2 + 1, current.maximum);
    }
    if (tree[id * 2].minimum < current.minimum) {
      apply_chmax(id * 2, current.minimum);
    }
    if (tree[id * 2 + 1].minimum < current.minimum) {
      apply_chmax(id * 2 + 1, current.minimum);
    }
  }

  void range_add_at(int id, int left, int right, int query_left,
                    int query_right, const T &value) {
    if (query_right <= left || right <= query_left) {
      return;
    }
    if (query_left <= left && right <= query_right) {
      apply_add(id, value);
      return;
    }
    push(id);
    int middle = (left + right) / 2;
    range_add_at(id * 2, left, middle, query_left, query_right, value);
    range_add_at(id * 2 + 1, middle, right, query_left, query_right, value);
    pull(id);
  }

  void range_chmin_at(int id, int left, int right, int query_left,
                      int query_right, const T &upper) {
    node &current = tree[id];
    if (query_right <= left || right <= query_left ||
        current.maximum <= upper) {
      return;
    }
    if (query_left <= left && right <= query_right &&
        (!current.has_second_maximum || current.second_maximum < upper)) {
      apply_chmin(id, upper);
      return;
    }
    push(id);
    int middle = (left + right) / 2;
    range_chmin_at(id * 2, left, middle, query_left, query_right, upper);
    range_chmin_at(id * 2 + 1, middle, right, query_left, query_right, upper);
    pull(id);
  }

  void range_chmax_at(int id, int left, int right, int query_left,
                      int query_right, const T &lower) {
    node &current = tree[id];
    if (query_right <= left || right <= query_left ||
        lower <= current.minimum) {
      return;
    }
    if (query_left <= left && right <= query_right &&
        (!current.has_second_minimum || lower < current.second_minimum)) {
      apply_chmax(id, lower);
      return;
    }
    push(id);
    int middle = (left + right) / 2;
    range_chmax_at(id * 2, left, middle, query_left, query_right, lower);
    range_chmax_at(id * 2 + 1, middle, right, query_left, query_right, lower);
    pull(id);
  }

  T range_sum_at(int id, int left, int right, int query_left,
                 int query_right) {
    if (query_left <= left && right <= query_right) {
      return tree[id].sum;
    }
    push(id);
    int middle = (left + right) / 2;
    T result{};
    if (query_left < middle) {
      result += range_sum_at(id * 2, left, middle, query_left, query_right);
    }
    if (middle < query_right) {
      result +=
          range_sum_at(id * 2 + 1, middle, right, query_left, query_right);
    }
    return result;
  }

  T range_min_at(int id, int left, int right, int query_left,
                 int query_right) {
    if (query_left <= left && right <= query_right) {
      return tree[id].minimum;
    }
    push(id);
    int middle = (left + right) / 2;
    T result = std::numeric_limits<T>::max();
    if (query_left < middle) {
      result = std::min(
          result, range_min_at(id * 2, left, middle, query_left, query_right));
    }
    if (middle < query_right) {
      result = std::min(result, range_min_at(id * 2 + 1, middle, right,
                                             query_left, query_right));
    }
    return result;
  }

  T range_max_at(int id, int left, int right, int query_left,
                 int query_right) {
    if (query_left <= left && right <= query_right) {
      return tree[id].maximum;
    }
    push(id);
    int middle = (left + right) / 2;
    T result = std::numeric_limits<T>::lowest();
    if (query_left < middle) {
      result = std::max(
          result, range_max_at(id * 2, left, middle, query_left, query_right));
    }
    if (middle < query_right) {
      result = std::max(result, range_max_at(id * 2 + 1, middle, right,
                                             query_left, query_right));
    }
    return result;
  }
};

} // namespace noya

#endif // NOYA_SEGMENT_TREE_BEATS_HPP
#include <algorithm>
#include <cassert>
#include <limits>
#include <vector>

/// @complexity Time: O(n) build; amortized O(log^2 n) chmin/chmax and O(log n) add/query.
/// Space: O(n).

namespace noya {

/// @brief Segment Tree Beats supporting range add, chmin, chmax, sum, minimum,
/// and maximum queries; bound updates run in amortized O(log^2 n) time.
template <class T> struct segment_tree_beats {
  struct node {
    T sum{};
    T maximum{};
    T second_maximum{};
    T minimum{};
    T second_minimum{};
    T lazy_add{};
    int maximum_count = 0;
    int minimum_count = 0;
    int length = 0;
    bool has_second_maximum = false;
    bool has_second_minimum = false;
  };

  int n = 0;
  std::vector<node> tree;

  segment_tree_beats() = default;
  explicit segment_tree_beats(const std::vector<T> &values) { build(values); }

  /// @brief Rebuild from values in O(n) time.
  void build(const std::vector<T> &values) {
    n = int(values.size());
    tree.assign(std::max(1, 4 * n), {});
    if (n > 0) {
      build_at(1, 0, n, values);
    }
  }

  /// @brief Add value to every element of [left, right).
  void range_add(int left, int right, const T &value) {
    check_range(left, right);
    if (left < right) {
      range_add_at(1, 0, n, left, right, value);
    }
  }

  /// @brief Replace each x in [left, right) by min(x, upper).
  void range_chmin(int left, int right, const T &upper) {
    check_range(left, right);
    if (left < right) {
      range_chmin_at(1, 0, n, left, right, upper);
    }
  }

  /// @brief Replace each x in [left, right) by max(x, lower).
  void range_chmax(int left, int right, const T &lower) {
    check_range(left, right);
    if (left < right) {
      range_chmax_at(1, 0, n, left, right, lower);
    }
  }

  /// @brief Return the sum over [left, right); an empty range has sum zero.
  T range_sum(int left, int right) {
    check_range(left, right);
    return left == right ? T{} : range_sum_at(1, 0, n, left, right);
  }

  /// @brief Return the minimum over the nonempty range [left, right).
  T range_min(int left, int right) {
    check_nonempty_range(left, right);
    return range_min_at(1, 0, n, left, right);
  }

  /// @brief Return the maximum over the nonempty range [left, right).
  T range_max(int left, int right) {
    check_nonempty_range(left, right);
    return range_max_at(1, 0, n, left, right);
  }

  /// @brief Return the sum of the entire array.
  T all_sum() const { return n == 0 ? T{} : tree[1].sum; }

  /// @brief Return the minimum of the nonempty array.
  T all_min() const {
    assert(n > 0);
    return tree[1].minimum;
  }

  /// @brief Return the maximum of the nonempty array.
  T all_max() const {
    assert(n > 0);
    return tree[1].maximum;
  }

private:
  void check_range(int left, int right) const {
    assert(0 <= left && left <= right && right <= n);
  }

  void check_nonempty_range(int left, int right) const {
    check_range(left, right);
    assert(left < right);
  }

  static node leaf(const T &value) {
    node result;
    result.sum = result.maximum = result.minimum = value;
    result.maximum_count = result.minimum_count = result.length = 1;
    return result;
  }

  static void assign_maximum(node &result, const node &left,
                             const node &right) {
    if (left.maximum == right.maximum) {
      result.maximum = left.maximum;
      result.maximum_count = left.maximum_count + right.maximum_count;
      if (left.has_second_maximum && right.has_second_maximum) {
        result.second_maximum =
            std::max(left.second_maximum, right.second_maximum);
        result.has_second_maximum = true;
      } else if (left.has_second_maximum) {
        result.second_maximum = left.second_maximum;
        result.has_second_maximum = true;
      } else if (right.has_second_maximum) {
        result.second_maximum = right.second_maximum;
        result.has_second_maximum = true;
      }
      return;
    }
    const node &higher = left.maximum > right.maximum ? left : right;
    const node &lower = left.maximum > right.maximum ? right : left;
    result.maximum = higher.maximum;
    result.maximum_count = higher.maximum_count;
    result.second_maximum =
        higher.has_second_maximum
            ? std::max(higher.second_maximum, lower.maximum)
            : lower.maximum;
    result.has_second_maximum = true;
  }

  static void assign_minimum(node &result, const node &left,
                             const node &right) {
    if (left.minimum == right.minimum) {
      result.minimum = left.minimum;
      result.minimum_count = left.minimum_count + right.minimum_count;
      if (left.has_second_minimum && right.has_second_minimum) {
        result.second_minimum =
            std::min(left.second_minimum, right.second_minimum);
        result.has_second_minimum = true;
      } else if (left.has_second_minimum) {
        result.second_minimum = left.second_minimum;
        result.has_second_minimum = true;
      } else if (right.has_second_minimum) {
        result.second_minimum = right.second_minimum;
        result.has_second_minimum = true;
      }
      return;
    }
    const node &lower = left.minimum < right.minimum ? left : right;
    const node &higher = left.minimum < right.minimum ? right : left;
    result.minimum = lower.minimum;
    result.minimum_count = lower.minimum_count;
    result.second_minimum =
        lower.has_second_minimum
            ? std::min(lower.second_minimum, higher.minimum)
            : higher.minimum;
    result.has_second_minimum = true;
  }

  static node merge_nodes(const node &left, const node &right) {
    node result;
    result.sum = left.sum + right.sum;
    result.length = left.length + right.length;
    assign_maximum(result, left, right);
    assign_minimum(result, left, right);
    return result;
  }

  void build_at(int id, int left, int right, const std::vector<T> &values) {
    if (right - left == 1) {
      tree[id] = leaf(values[left]);
      return;
    }
    int middle = (left + right) / 2;
    build_at(id * 2, left, middle, values);
    build_at(id * 2 + 1, middle, right, values);
    pull(id);
  }

  void pull(int id) { tree[id] = merge_nodes(tree[id * 2], tree[id * 2 + 1]); }

  void apply_add(int id, const T &value) {
    node &current = tree[id];
    current.sum += value * current.length;
    current.maximum += value;
    current.minimum += value;
    if (current.has_second_maximum) {
      current.second_maximum += value;
    }
    if (current.has_second_minimum) {
      current.second_minimum += value;
    }
    current.lazy_add += value;
  }

  void apply_chmin(int id, const T &upper) {
    node &current = tree[id];
    if (current.maximum <= upper) {
      return;
    }
    T old_maximum = current.maximum;
    current.sum += (upper - old_maximum) * current.maximum_count;
    if (current.minimum == old_maximum) {
      current.minimum = upper;
    } else if (current.has_second_minimum &&
               current.second_minimum == old_maximum) {
      current.second_minimum = upper;
    }
    current.maximum = upper;
  }

  void apply_chmax(int id, const T &lower) {
    node &current = tree[id];
    if (lower <= current.minimum) {
      return;
    }
    T old_minimum = current.minimum;
    current.sum += (lower - old_minimum) * current.minimum_count;
    if (current.maximum == old_minimum) {
      current.maximum = lower;
    } else if (current.has_second_maximum &&
               current.second_maximum == old_minimum) {
      current.second_maximum = lower;
    }
    current.minimum = lower;
  }

  void push(int id) {
    node &current = tree[id];
    if (current.length == 1) {
      current.lazy_add = T{};
      return;
    }
    if (current.lazy_add != T{}) {
      apply_add(id * 2, current.lazy_add);
      apply_add(id * 2 + 1, current.lazy_add);
      current.lazy_add = T{};
    }
    if (tree[id * 2].maximum > current.maximum) {
      apply_chmin(id * 2, current.maximum);
    }
    if (tree[id * 2 + 1].maximum > current.maximum) {
      apply_chmin(id * 2 + 1, current.maximum);
    }
    if (tree[id * 2].minimum < current.minimum) {
      apply_chmax(id * 2, current.minimum);
    }
    if (tree[id * 2 + 1].minimum < current.minimum) {
      apply_chmax(id * 2 + 1, current.minimum);
    }
  }

  void range_add_at(int id, int left, int right, int query_left,
                    int query_right, const T &value) {
    if (query_right <= left || right <= query_left) {
      return;
    }
    if (query_left <= left && right <= query_right) {
      apply_add(id, value);
      return;
    }
    push(id);
    int middle = (left + right) / 2;
    range_add_at(id * 2, left, middle, query_left, query_right, value);
    range_add_at(id * 2 + 1, middle, right, query_left, query_right, value);
    pull(id);
  }

  void range_chmin_at(int id, int left, int right, int query_left,
                      int query_right, const T &upper) {
    node &current = tree[id];
    if (query_right <= left || right <= query_left ||
        current.maximum <= upper) {
      return;
    }
    if (query_left <= left && right <= query_right &&
        (!current.has_second_maximum || current.second_maximum < upper)) {
      apply_chmin(id, upper);
      return;
    }
    push(id);
    int middle = (left + right) / 2;
    range_chmin_at(id * 2, left, middle, query_left, query_right, upper);
    range_chmin_at(id * 2 + 1, middle, right, query_left, query_right, upper);
    pull(id);
  }

  void range_chmax_at(int id, int left, int right, int query_left,
                      int query_right, const T &lower) {
    node &current = tree[id];
    if (query_right <= left || right <= query_left ||
        lower <= current.minimum) {
      return;
    }
    if (query_left <= left && right <= query_right &&
        (!current.has_second_minimum || lower < current.second_minimum)) {
      apply_chmax(id, lower);
      return;
    }
    push(id);
    int middle = (left + right) / 2;
    range_chmax_at(id * 2, left, middle, query_left, query_right, lower);
    range_chmax_at(id * 2 + 1, middle, right, query_left, query_right, lower);
    pull(id);
  }

  T range_sum_at(int id, int left, int right, int query_left,
                 int query_right) {
    if (query_left <= left && right <= query_right) {
      return tree[id].sum;
    }
    push(id);
    int middle = (left + right) / 2;
    T result{};
    if (query_left < middle) {
      result += range_sum_at(id * 2, left, middle, query_left, query_right);
    }
    if (middle < query_right) {
      result +=
          range_sum_at(id * 2 + 1, middle, right, query_left, query_right);
    }
    return result;
  }

  T range_min_at(int id, int left, int right, int query_left,
                 int query_right) {
    if (query_left <= left && right <= query_right) {
      return tree[id].minimum;
    }
    push(id);
    int middle = (left + right) / 2;
    T result = std::numeric_limits<T>::max();
    if (query_left < middle) {
      result = std::min(
          result, range_min_at(id * 2, left, middle, query_left, query_right));
    }
    if (middle < query_right) {
      result = std::min(result, range_min_at(id * 2 + 1, middle, right,
                                             query_left, query_right));
    }
    return result;
  }

  T range_max_at(int id, int left, int right, int query_left,
                 int query_right) {
    if (query_left <= left && right <= query_right) {
      return tree[id].maximum;
    }
    push(id);
    int middle = (left + right) / 2;
    T result = std::numeric_limits<T>::lowest();
    if (query_left < middle) {
      result = std::max(
          result, range_max_at(id * 2, left, middle, query_left, query_right));
    }
    if (middle < query_right) {
      result = std::max(result, range_max_at(id * 2 + 1, middle, right,
                                             query_left, query_right));
    }
    return result;
  }
};

} // namespace noya