把路径上的点权最短路压成 k<=3 的 min-plus 矩阵,并用重链剖分维护路径转移。
OJ: luogu
题目 ID: P8820
难度:提高+/省选-
标签:图论最短路树形结构倍增矩阵
日期: 2026-07-06 08:46
题意
给定一棵 n 个点的树。两台主机之间的树上距离不超过 k 时,可以直接传输数据。一次请求从 s 到 t,可以选择若干中转主机,要求相邻中转主机之间都能直接传输。
经过的每台主机都要付出点权 v_i,包括起点和终点。要求每次询问的最小总代价。
思路
小数据可以先求出任意两点树上距离,再把所有距离不超过 k 的点对连边,最后在这个新图上求点权最短路:
// brute.cpp:小数据暴力解,先 Floyd 求树上距离,再在可直接传输图上 Floyd 求最短路。
#include <bits/stdc++.h>
using namespace std;
const long long INF = (long long)4e18;
const int MAXN = 55;
int n, q, K;
long long value_cost[MAXN];
long long tree_dist[MAXN][MAXN];
long long answer_dist[MAXN][MAXN];
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
cin >> n >> q >> K;
for (int i = 1; i <= n; i++) {
cin >> value_cost[i];
}
for (int i = 1; i <= n; i++) {
for (int j = 1; j <= n; j++) {
tree_dist[i][j] = (i == j) ? 0 : INF;
}
}
for (int i = 1; i < n; i++) {
int u, v;
cin >> u >> v;
tree_dist[u][v] = tree_dist[v][u] = 1;
}
for (int k = 1; k <= n; k++) {
for (int i = 1; i <= n; i++) {
for (int j = 1; j <= n; j++) {
if (tree_dist[i][j] > tree_dist[i][k] + tree_dist[k][j]) {
tree_dist[i][j] = tree_dist[i][k] + tree_dist[k][j];
}
}
}
}
for (int i = 1; i <= n; i++) {
for (int j = 1; j <= n; j++) {
answer_dist[i][j] = INF;
}
answer_dist[i][i] = value_cost[i];
}
for (int i = 1; i <= n; i++) {
for (int j = 1; j <= n; j++) {
if (i != j && tree_dist[i][j] <= K) {
answer_dist[i][j] = value_cost[i] + value_cost[j];
}
}
}
for (int k = 1; k <= n; k++) {
for (int i = 1; i <= n; i++) {
for (int j = 1; j <= n; j++) {
if (answer_dist[i][j] > answer_dist[i][k] + answer_dist[k][j] - value_cost[k]) {
answer_dist[i][j] = answer_dist[i][k] + answer_dist[k][j] - value_cost[k];
}
}
}
}
while (q--) {
int s, t;
cin >> s >> t;
cout << answer_dist[s][t] << '\n';
}
return 0;
}这个建模非常直接:把每台主机看成一个状态。如果当前数据在主机 u 上,那么下一步可以传到所有树上距离不超过 k 的主机 v,并额外付出 v_v 的处理时间。
因此从 s 出发时初始代价是 v_s,每条转移 u -> v 的代价是 v_v。这就是普通的最短路问题。
满数据的难点在于:不能为每次询问重新建图或跑 Dijkstra。注意 k <= 3,而一次询问只关心树上 s 到 t 的路径。最优传输路线可以理解为:沿着这条路径从两端往中间走,允许一次跨过不超过 k 条边。
于是我们维护一个长度为 k 的 DP 状态:
dp[d]:当前已经处理到某个路径点,距离上一次被选作中转主机的点有 d 条边时的最小代价因为 d 只可能是 0..k-1,每向父亲方向走过一个点,都可以用一个 k*k 的 min-plus 矩阵表示状态转移:
- 选择这个点作为中转点:状态回到
0,代价加上这个点的点权; - 不选择这个点:距离
d增加1; - 当
k=3时,还要处理从路径两侧同时离路径一步、通过某个邻接点连接的特殊情况,代码中用min_neighbor[u]记录u的相邻点最小点权。
这样一段路径的转移就是若干矩阵的 min-plus 乘积。
为了快速取得从某个点向上走到链顶的转移矩阵,代码使用重链剖分:
base_matrix[u]表示从u走到父亲时的一步转移;chain_matrix[u]表示从u一直走到当前重链链顶父亲的转移;- 线段树按 DFS 序维护重链内部的矩阵乘积,用来处理同一条重链上的中间一段。
查询 s,t 时,分别从两端向 LCA 收缩,得到两侧 DP 状态 left_state 和 right_state。最后枚举两侧末端距离 i,j,只要 i+j <= k,两边就可以接起来,取最小代价。
代码
// main.cpp:k<=3 的树上点权最短路,用重链剖分维护 min-plus 转移矩阵。
#include <bits/stdc++.h>
using namespace std;
const int MAXN = 200005;
const int MAXM = 400005;
const long long INF = (long long)4e18;
struct Matrix {
long long a[3][3];
};
struct DpState {
long long a[3];
};
int n, q, K;
long long val[MAXN], min_neighbor[MAXN];
int head[MAXN], to[MAXM], nxt[MAXM], edge_cnt;
int parent_node[MAXN], depth_node[MAXN], subtree_size[MAXN], heavy_son[MAXN];
int top_node[MAXN], dfn[MAXN], rev_dfn[MAXN], dfn_cnt;
Matrix base_matrix[MAXN], chain_matrix[MAXN], seg_tree[MAXN * 4];
void add_edge(int u, int v) {
edge_cnt++;
to[edge_cnt] = v;
nxt[edge_cnt] = head[u];
head[u] = edge_cnt;
}
long long safe_add(long long x, long long y) {
if (x >= INF / 2 || y >= INF / 2) {
return INF;
}
if (x + y >= INF) {
return INF;
}
return x + y;
}
Matrix multiply_matrix(const Matrix &x, const Matrix &y) {
Matrix result;
for (int i = 0; i < 3; i++) {
for (int j = 0; j < 3; j++) {
result.a[i][j] = INF;
}
}
for (int i = 0; i < K; i++) {
for (int j = 0; j < K; j++) {
for (int k = 0; k < K; k++) {
result.a[i][j] = min(result.a[i][j], safe_add(x.a[i][k], y.a[k][j]));
}
}
}
return result;
}
DpState multiply_dp(const DpState &x, const Matrix &y) {
DpState result;
for (int i = 0; i < 3; i++) {
result.a[i] = INF;
}
for (int i = 0; i < K; i++) {
for (int j = 0; j < K; j++) {
result.a[i] = min(result.a[i], safe_add(x.a[j], y.a[j][i]));
}
}
return result;
}
Matrix make_transition(long long x, long long mn) {
Matrix result;
for (int i = 0; i < 3; i++) {
for (int j = 0; j < 3; j++) {
result.a[i][j] = INF;
}
}
if (K == 1) {
result.a[0][0] = x;
} else if (K == 2) {
result.a[0][0] = x;
result.a[1][0] = x;
result.a[0][1] = 0;
} else {
result.a[0][0] = x;
result.a[1][0] = x;
result.a[2][0] = x;
result.a[0][1] = 0;
result.a[1][2] = 0;
result.a[2][2] = mn;
}
return result;
}
void build_tree_info() {
vector<int> order;
order.reserve(n);
stack<int> st;
st.push(1);
parent_node[1] = 0;
depth_node[1] = 1;
while (!st.empty()) {
int u = st.top();
st.pop();
order.push_back(u);
for (int e = head[u]; e != 0; e = nxt[e]) {
int v = to[e];
if (v == parent_node[u]) {
continue;
}
parent_node[v] = u;
depth_node[v] = depth_node[u] + 1;
st.push(v);
}
}
for (int i = 1; i <= n; i++) {
min_neighbor[i] = INF;
}
for (int u = 1; u <= n; u++) {
for (int e = head[u]; e != 0; e = nxt[e]) {
int v = to[e];
min_neighbor[u] = min(min_neighbor[u], val[v]);
}
}
for (int i = (int)order.size() - 1; i >= 0; i--) {
int u = order[i];
subtree_size[u] = 1;
heavy_son[u] = 0;
for (int e = head[u]; e != 0; e = nxt[e]) {
int v = to[e];
if (v == parent_node[u]) {
continue;
}
subtree_size[u] += subtree_size[v];
if (subtree_size[v] > subtree_size[heavy_son[u]]) {
heavy_son[u] = v;
}
}
}
stack<pair<int, int> > starts;
starts.push(make_pair(1, 1));
while (!starts.empty()) {
int start = starts.top().first;
int top = starts.top().second;
starts.pop();
int u = start;
while (u != 0) {
top_node[u] = top;
dfn[u] = ++dfn_cnt;
rev_dfn[dfn_cnt] = u;
for (int e = head[u]; e != 0; e = nxt[e]) {
int v = to[e];
if (v == parent_node[u] || v == heavy_son[u]) {
continue;
}
starts.push(make_pair(v, v));
}
u = heavy_son[u];
}
}
for (int i = 1; i <= n; i++) {
int u = order[i - 1];
long long parent_value = (parent_node[u] == 0) ? INF : val[parent_node[u]];
base_matrix[u] = make_transition(parent_value, min_neighbor[u]);
if (u == top_node[u]) {
chain_matrix[u] = base_matrix[u];
} else {
chain_matrix[u] = multiply_matrix(base_matrix[u], chain_matrix[parent_node[u]]);
}
}
}
void build_segment_tree(int node, int l, int r) {
if (l == r) {
seg_tree[node] = base_matrix[rev_dfn[l]];
return;
}
int mid = (l + r) / 2;
build_segment_tree(node * 2, l, mid);
build_segment_tree(node * 2 + 1, mid + 1, r);
seg_tree[node] = multiply_matrix(seg_tree[node * 2 + 1], seg_tree[node * 2]);
}
Matrix query_segment_tree(int ql, int qr, int node, int l, int r) {
if (ql <= l && r <= qr) {
return seg_tree[node];
}
int mid = (l + r) / 2;
if (qr <= mid) {
return query_segment_tree(ql, qr, node * 2, l, mid);
}
if (ql > mid) {
return query_segment_tree(ql, qr, node * 2 + 1, mid + 1, r);
}
Matrix right_part = query_segment_tree(ql, qr, node * 2 + 1, mid + 1, r);
Matrix left_part = query_segment_tree(ql, qr, node * 2, l, mid);
return multiply_matrix(right_part, left_part);
}
long long solve_query(int u, int v) {
if (u == v) {
return val[u];
}
DpState left_state, right_state;
for (int i = 0; i < 3; i++) {
left_state.a[i] = right_state.a[i] = INF;
}
left_state.a[0] = val[u];
right_state.a[0] = val[v];
while (top_node[u] != top_node[v]) {
if (depth_node[top_node[u]] < depth_node[top_node[v]]) {
swap(u, v);
swap(left_state, right_state);
}
left_state = multiply_dp(left_state, chain_matrix[u]);
u = parent_node[top_node[u]];
}
if (depth_node[u] > depth_node[v]) {
swap(u, v);
swap(left_state, right_state);
}
if (u != v) {
Matrix middle = query_segment_tree(dfn[u] + 1, dfn[v], 1, 1, n);
right_state = multiply_dp(right_state, middle);
}
long long answer = left_state.a[0] + right_state.a[0] - val[u];
for (int i = 0; i < K; i++) {
for (int j = 0; j < K; j++) {
if (i == 0 && j == 0) {
continue;
}
if (i + j <= K) {
answer = min(answer, safe_add(left_state.a[i], right_state.a[j]));
}
}
}
if (K == 3) {
answer = min(answer, safe_add(safe_add(left_state.a[2], right_state.a[2]), min_neighbor[u]));
}
return answer;
}
int main() {
ios::sync_with_stdio(false);
cin.tie(nullptr);
cin >> n >> q >> K;
for (int i = 1; i <= n; i++) {
cin >> val[i];
}
for (int i = 1; i < n; i++) {
int u, v;
cin >> u >> v;
add_edge(u, v);
add_edge(v, u);
}
build_tree_info();
build_segment_tree(1, 1, n);
while (q--) {
int u, v;
cin >> u >> v;
cout << solve_query(u, v) << '\n';
}
return 0;
}复杂度
预处理重链剖分和线段树为 k <= 3,可以看成
每次询问会跳过若干条重链,每次做常数大小矩阵或 DP 转移,复杂度为
空间复杂度为
总结
本题最基础的模型是“树的 k 次幂图上的点权最短路”。满分做法的关键,是发现 k <= 3 让“走过一段路径”可以压成很小的 min-plus 矩阵。
重链剖分负责把任意树上路径拆成少量连续链段,矩阵乘法负责把每段链上的选择/跳过决策合并起来。