用树链剖分把树上路径拆成若干个 DFS 序区间,再在线段树中同时维护区间和与区间最大值。
OJ: luogu
题目 ID: P2590
难度:普及+/提高
标签:树链剖分线段树dfs序树区间最大值
日期: 2026-06-21 03:09
题意
给出一棵带点权的树,支持三种操作:
- 单点修改权值
- 查询两点路径上的最大点权
- 查询两点路径上的点权和
思路
先看一个可以直接验证想法的朴素解:
cpp
#include <bits/stdc++.h>
using namespace std;
const int MAXN = 30000 + 5;
int n;
vector<int> g[MAXN];
int val[MAXN];
int parent_arr[MAXN];
int depth_arr[MAXN];
void build_parent() {
queue<int> q;
q.push(1);
parent_arr[1] = 0;
depth_arr[1] = 1;
while (!q.empty()) {
int u = q.front();
q.pop();
for (int v : g[u]) {
if (v == parent_arr[u]) {
continue;
}
parent_arr[v] = u;
depth_arr[v] = depth_arr[u] + 1;
q.push(v);
}
}
}
long long query_path_sum(int u, int v) {
long long ans = 0;
while (depth_arr[u] > depth_arr[v]) {
ans += val[u];
u = parent_arr[u];
}
while (depth_arr[v] > depth_arr[u]) {
ans += val[v];
v = parent_arr[v];
}
while (u != v) {
ans += val[u] + val[v];
u = parent_arr[u];
v = parent_arr[v];
}
ans += val[u];
return ans;
}
int query_path_max(int u, int v) {
int ans = -30000;
while (depth_arr[u] > depth_arr[v]) {
ans = max(ans, val[u]);
u = parent_arr[u];
}
while (depth_arr[v] > depth_arr[u]) {
ans = max(ans, val[v]);
v = parent_arr[v];
}
while (u != v) {
ans = max(ans, val[u]);
ans = max(ans, val[v]);
u = parent_arr[u];
v = parent_arr[v];
}
ans = max(ans, val[u]);
return ans;
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
// 这是一个朴素做法:
// 每次查询都直接沿父亲往上跳,适合小数据对拍。
cin >> n;
for (int i = 1; i < n; i++) {
int u, v;
cin >> u >> v;
g[u].push_back(v);
g[v].push_back(u);
}
for (int i = 1; i <= n; i++) {
cin >> val[i];
}
build_parent();
int q;
cin >> q;
while (q--) {
string op;
int u, v;
cin >> op >> u >> v;
if (op[1] == 'H') {
val[u] = v;
} else if (op[1] == 'M') {
cout << query_path_max(u, v) << '\n';
} else {
cout << query_path_sum(u, v) << '\n';
}
}
return 0;
}brute.cpp 直接沿父亲链往上跳,直到两点相遇。
这个方法完全正确,但如果树退化成链,一次查询可能要走
这题的核心是把树上路径转成区间问题。
树链剖分后:
- 每个点有一个
dfn - 同一条重链在 DFS 序里是连续区间
于是任意路径都能拆成若干个连续区间。
接下来再用线段树维护这些区间的信息即可。
因为题目同时需要:
- 路径和
- 路径最大值
所以在线段树每个节点里同时维护:
- 区间和
- 区间最大值
CHANGE 做单点修改。
QSUM 查询路径时,把各段区间和累加起来。
QMAX 查询路径时,把各段区间最大值再取最大。
要特别注意一点:
这题点权可能为负数,所以路径最大值查询的初值不能写成 0,必须是一个足够小的负数。
代码
cpp
#include <bits/stdc++.h>
using namespace std;
const int MAXN = 30000 + 5;
const int NEG_INF = -0x3f3f3f3f;
int n;
int head[MAXN], to[MAXN << 1], nxt[MAXN << 1], edge_cnt;
int val[MAXN];
int parent_arr[MAXN], depth_arr[MAXN];
int sub_size[MAXN], heavy_son[MAXN];
int top_arr[MAXN], dfn[MAXN], rev_dfn[MAXN], dfs_clock;
long long seg_sum[MAXN << 2];
int seg_max[MAXN << 2];
void add_edge(int u, int v) {
edge_cnt++;
to[edge_cnt] = v;
nxt[edge_cnt] = head[u];
head[u] = edge_cnt;
}
void build_tree_info() {
static int order[MAXN];
static int stk[MAXN];
int ord_cnt = 0;
int top = 0;
stk[++top] = 1;
parent_arr[1] = 0;
depth_arr[1] = 1;
while (top > 0) {
int u = stk[top--];
order[++ord_cnt] = u;
for (int i = head[u]; i != 0; i = nxt[i]) {
int v = to[i];
if (v == parent_arr[u]) {
continue;
}
parent_arr[v] = u;
depth_arr[v] = depth_arr[u] + 1;
stk[++top] = v;
}
}
for (int idx = ord_cnt; idx >= 1; idx--) {
int u = order[idx];
sub_size[u] = 1;
heavy_son[u] = -1;
int best_size = 0;
for (int i = head[u]; i != 0; i = nxt[i]) {
int v = to[i];
if (v == parent_arr[u]) {
continue;
}
sub_size[u] += sub_size[v];
if (sub_size[v] > best_size) {
best_size = sub_size[v];
heavy_son[u] = v;
}
}
}
}
void build_dfn() {
static int stk_u[MAXN];
static int stk_top[MAXN];
int top = 0;
stk_u[++top] = 1;
stk_top[top] = 1;
while (top > 0) {
int u = stk_u[top];
int chain_top = stk_top[top];
top--;
while (u != -1) {
top_arr[u] = chain_top;
dfn[u] = ++dfs_clock;
rev_dfn[dfs_clock] = u;
for (int i = head[u]; i != 0; i = nxt[i]) {
int v = to[i];
if (v == parent_arr[u] || v == heavy_son[u]) {
continue;
}
stk_u[++top] = v;
stk_top[top] = v;
}
u = heavy_son[u];
}
}
}
void push_up(int u) {
seg_sum[u] = seg_sum[u << 1] + seg_sum[u << 1 | 1];
seg_max[u] = max(seg_max[u << 1], seg_max[u << 1 | 1]);
}
void build_seg(int u, int l, int r) {
if (l == r) {
int node = rev_dfn[l];
seg_sum[u] = val[node];
seg_max[u] = val[node];
return;
}
int mid = (l + r) >> 1;
build_seg(u << 1, l, mid);
build_seg(u << 1 | 1, mid + 1, r);
push_up(u);
}
void point_update(int u, int l, int r, int pos, int new_val) {
if (l == r) {
seg_sum[u] = new_val;
seg_max[u] = new_val;
return;
}
int mid = (l + r) >> 1;
if (pos <= mid) {
point_update(u << 1, l, mid, pos, new_val);
} else {
point_update(u << 1 | 1, mid + 1, r, pos, new_val);
}
push_up(u);
}
long long query_sum(int u, int l, int r, int ql, int qr) {
if (ql <= l && r <= qr) {
return seg_sum[u];
}
int mid = (l + r) >> 1;
long long ans = 0;
if (ql <= mid) {
ans += query_sum(u << 1, l, mid, ql, qr);
}
if (qr > mid) {
ans += query_sum(u << 1 | 1, mid + 1, r, ql, qr);
}
return ans;
}
int query_max(int u, int l, int r, int ql, int qr) {
if (ql <= l && r <= qr) {
return seg_max[u];
}
int mid = (l + r) >> 1;
int ans = NEG_INF;
if (ql <= mid) {
ans = max(ans, query_max(u << 1, l, mid, ql, qr));
}
if (qr > mid) {
ans = max(ans, query_max(u << 1 | 1, mid + 1, r, ql, qr));
}
return ans;
}
long long query_path_sum(int u, int v) {
long long ans = 0;
while (top_arr[u] != top_arr[v]) {
if (depth_arr[top_arr[u]] < depth_arr[top_arr[v]]) {
swap(u, v);
}
ans += query_sum(1, 1, n, dfn[top_arr[u]], dfn[u]);
u = parent_arr[top_arr[u]];
}
if (depth_arr[u] > depth_arr[v]) {
swap(u, v);
}
ans += query_sum(1, 1, n, dfn[u], dfn[v]);
return ans;
}
int query_path_max(int u, int v) {
int ans = NEG_INF;
while (top_arr[u] != top_arr[v]) {
if (depth_arr[top_arr[u]] < depth_arr[top_arr[v]]) {
swap(u, v);
}
ans = max(ans, query_max(1, 1, n, dfn[top_arr[u]], dfn[u]));
u = parent_arr[top_arr[u]];
}
if (depth_arr[u] > depth_arr[v]) {
swap(u, v);
}
ans = max(ans, query_max(1, 1, n, dfn[u], dfn[v]));
return ans;
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
cin >> n;
for (int i = 1; i < n; i++) {
int u, v;
cin >> u >> v;
add_edge(u, v);
add_edge(v, u);
}
for (int i = 1; i <= n; i++) {
cin >> val[i];
}
build_tree_info();
build_dfn();
build_seg(1, 1, n);
int q;
cin >> q;
while (q--) {
string op;
int u, v;
cin >> op >> u >> v;
if (op[1] == 'H') {
val[u] = v;
point_update(1, 1, n, dfn[u], v);
} else if (op[1] == 'M') {
cout << query_path_max(u, v) << '\n';
} else {
cout << query_path_sum(u, v) << '\n';
}
}
return 0;
}复杂度
预处理和建树是
单点修改是
每次路径查询会被拆成 QSUM 和 QMAX 的复杂度都是
空间复杂度是
总结
这题是很标准的树链剖分模板题。
关键是两步:
- 先把树上路径拆成若干个 DFS 序连续区间
- 再在线段树上维护这些区间的统计信息
而且要记住负权这个细节,否则 QMAX 很容易写错。