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¶
#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 ¤t = 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 ¤t = 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 ¤t = 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 ¤t = 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 ¤t = 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 ¤t = 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 ¤t = 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 ¤t = 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 ¤t = 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 ¤t = 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 ¤t = 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 ¤t = 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