线段树详解:区间修改与查询的高效实现
区间加法与区间求和
线段树是一种用于维护区间信息的数据结构,可以在 $O(\log n)$ 的时间复杂度内完成单点或区间修改以及区间查询操作。以下是针对区间加法和区间求和的三种实现方式。
1. 模板类封装版
使用 C++ 模板类封装,支持多种数据类型,逻辑清晰,适合工程化应用。
#include <iostream>
#include <vector>
#include <algorithm>
template <typename T>
class RangeSumTree {
private:
int size;
std::vector<T> nodes;
std::vector<T> pending_add;
void propagate(int idx, int s, int e) {
if (pending_add[idx] == 0) return;
int mid = (s + e) >> 1;
int left = idx << 1, right = idx << 1 | 1;
nodes[left] += pending_add[idx] * (mid - s + 1);
pending_add[left] += pending_add[idx];
nodes[right] += pending_add[idx] * (e - mid);
pending_add[right] += pending_add[idx];
pending_add[idx] = 0;
}
void build_tree(int idx, int s, int e, const std::vector<T>& initial_data) {
if (s == e) {
nodes[idx] = initial_data[s];
return;
}
int mid = (s + e) >> 1;
build_tree(idx << 1, s, mid, initial_data);
build_tree(idx << 1 | 1, mid + 1, e, initial_data);
nodes[idx] = nodes[idx << 1] + nodes[idx << 1 | 1];
}
void modify(int idx, int s, int e, int l, int r, T delta) {
if (l <= s && e <= r) {
nodes[idx] += delta * (e - s + 1);
pending_add[idx] += delta;
return;
}
propagate(idx, s, e);
int mid = (s + e) >> 1;
if (l <= mid) modify(idx << 1, s, mid, l, r, delta);
if (r > mid) modify(idx << 1 | 1, mid + 1, e, l, r, delta);
nodes[idx] = nodes[idx << 1] + nodes[idx << 1 | 1];
}
T get_sum(int idx, int s, int e, int l, int r) {
if (l <= s && e <= r) return nodes[idx];
propagate(idx, s, e);
int mid = (s + e) >> 1;
T res = 0;
if (l <= mid) res += get_sum(idx << 1, s, mid, l, r);
if (r > mid) res += get_sum(idx << 1 | 1, mid + 1, e, l, r);
return res;
}
public:
RangeSumTree(const std::vector<T>& data) {
size = data.size() - 1;
nodes.assign(size * 4, 0);
pending_add.assign(size * 4, 0);
build_tree(1, 1, size, data);
}
void add(int l, int r, T val) {
if (l > r) std::swap(l, r);
modify(1, 1, size, l, r, val);
}
T query(int l, int r) {
if (l > r) std::swap(l, r);
return get_sum(1, 1, size, l, r);
}
};
2. 结构体数组版
将节点信息封装在结构体中,逻辑直观,适合比赛快速书写。
const int MAXN = 1e5 + 7;
struct Node {
int left, right;
long long val, tag;
} tree[MAXN * 4];
long long raw_data[MAXN];
void update_node(int u) {
tree[u].val = tree[u << 1].val + tree[u << 1 | 1].val;
}
void spread_tag(int u) {
if (tree[u].tag) {
Node &l = tree[u << 1], &r = tree[u << 1 | 1];
l.val += tree[u].tag * (l.right - l.left + 1);
l.tag += tree[u].tag;
r.val += tree[u].tag * (r.right - r.left + 1);
r.tag += tree[u].tag;
tree[u].tag = 0;
}
}
void build(int u, int l, int r) {
tree[u] = {l, r, 0, 0};
if (l == r) {
tree[u].val = raw_data[l];
return;
}
int mid = (l + r) >> 1;
build(u << 1, l, mid);
build(u << 1 | 1, mid + 1, r);
update_node(u);
}
void range_add(int u, int l, int r, long long k) {
if (l <= tree[u].left && tree[u].right <= r) {
tree[u].val += k * (tree[u].right - tree[u].left + 1);
tree[u].tag += k;
return;
}
spread_tag(u);
int mid = (tree[u].left + tree[u].right) >> 1;
if (l <= mid) range_add(u << 1, l, r, k);
if (r > mid) range_add(u << 1 | 1, l, r, k);
update_node(u);
}
long long range_query(int u, int l, int r) {
if (l <= tree[u].left && tree[u].right <= r) return tree[u].val;
spread_tag(u);
int mid = (tree[u].left + tree[u].right) >> 1;
long long total = 0;
if (l <= mid) total += range_query(u << 1, l, r);
if (r > mid) total += range_query(u << 1 | 1, l, r);
return total;
}
区间加法与区间最大值
在维护区间最值时,懒标记的下传不需要乘以区间长度,因为区间内每个数都增加 $k$ 时,最大值也仅仅增加 $k$。
template <typename T>
class RangeMaxTree {
private:
int n;
std::vector<T> tree_max;
std::vector<T> lazy_val;
void push_down(int p) {
if (lazy_val[p] == 0) return;
tree_max[p << 1] += lazy_val[p];
lazy_val[p << 1] += lazy_val[p];
tree_max[p << 1 | 1] += lazy_val[p];
lazy_val[p << 1 | 1] += lazy_val[p];
lazy_val[p] = 0;
}
void update(int p, int s, int e, int l, int r, T v) {
if (l <= s && e <= r) {
tree_max[p] += v;
lazy_val[p] += v;
return;
}
push_down(p);
int mid = (s + e) >> 1;
if (l <= mid) update(p << 1, s, mid, l, r, v);
if (r > mid) update(p << 1 | 1, mid + 1, e, l, r, v);
tree_max[p] = std::max(tree_max[p << 1], tree_max[p << 1 | 1]);
}
T query(int p, int s, int e, int l, int r) {
if (l <= s && e <= r) return tree_max[p];
push_down(p);
int mid = (s + e) >> 1;
T res = std::numeric_limits<T>::lowest();
if (l <= mid) res = std::max(res, query(p << 1, s, mid, l, r));
if (r > mid) res = std::max(res, query(p << 1 | 1, mid + 1, e, l, r));
return res;
}
public:
RangeMaxTree(const std::vector<T>& src) : n(src.size() - 1) {
tree_max.assign(4 * n, 0);
lazy_val.assign(4 * n, 0);
auto init = [&](auto self, int p, int s, int e) -> void {
if (s == e) {
tree_max[p] = src[s];
return;
}
int mid = (s + e) >> 1;
self(self, p << 1, s, mid);
self(self, p << 1 | 1, mid + 1, e);
tree_max[p] = std::max(tree_max[p << 1], tree_max[p << 1 | 1]);
};
init(init, 1, 1, n);
}
void range_add(int l, int r, T v) { update(1, 1, n, l, r, v); }
T range_max(int l, int r) { return query(1, 1, n, l, r); }
};
区间加法、区间乘法与区间求和
当同时存在加法和乘法操作时,需要维护两个懒标记。约定标记下传的顺序:先乘后加。即 $val = val \times mul\_tag + add\_tag$。
template <typename T>
class AdvancedSegmentTree {
private:
int n;
T MOD;
std::vector<T> sum, mul, add;
void apply(int u, int l, int r, T m, T a) {
sum[u] = (sum[u] * m + a * (r - l + 1)) % MOD;
mul[u] = (mul[u] * m) % MOD;
add[u] = (add[u] * m + a) % MOD;
}
void push_down(int u, int l, int r) {
int mid = (l + r) >> 1;
apply(u << 1, l, mid, mul[u], add[u]);
apply(u << 1 | 1, mid + 1, r, mul[u], add[u]);
mul[u] = 1;
add[u] = 0;
}
void build(int u, int l, int r, const std::vector<T>& a) {
mul[u] = 1;
if (l == r) {
sum[u] = a[l] % MOD;
return;
}
int mid = (l + r) >> 1;
build(u << 1, l, mid, a);
build(u << 1 | 1, mid + 1, r, a);
sum[u] = (sum[u << 1] + sum[u << 1 | 1]) % MOD;
}
void modify(int u, int l, int r, int ql, int qr, T m, T a) {
if (ql <= l && r <= qr) {
apply(u, l, r, m, a);
return;
}
push_down(u, l, r);
int mid = (l + r) >> 1;
if (ql <= mid) modify(u << 1, l, mid, ql, qr, m, a);
if (qr > mid) modify(u << 1 | 1, mid + 1, r, ql, qr, m, a);
sum[u] = (sum[u << 1] + sum[u << 1 | 1]) % MOD;
}
T ask(int u, int l, int r, int ql, int qr) {
if (ql <= l && r <= qr) return sum[u];
push_down(u, l, r);
int mid = (l + r) >> 1;
T res = 0;
if (ql <= mid) res = (res + ask(u << 1, l, mid, ql, qr)) % MOD;
if (qr > mid) res = (res + ask(u << 1 | 1, mid + 1, r, ql, qr)) % MOD;
return res;
}
public:
AdvancedSegmentTree(const std::vector<T>& data, T m_val) : MOD(m_val) {
n = data.size() - 1;
sum.resize(n * 4);
mul.assign(n * 4, 1);
add.assign(n * 4, 0);
build(1, 1, n, data);
}
void multiply(int l, int r, T k) { modify(1, 1, n, l, r, k % MOD, 0); }
void add_val(int l, int r, T k) { modify(1, 1, n, l, r, 1, k % MOD); }
T query_sum(int l, int r) { return ask(1, 1, n, l, r); }
};