[CSP-S 2022] 数据传输

GitHub跳转原题关系图返回列表

把路径上的点权最短路压成 k<=3 的 min-plus 矩阵,并用重链剖分维护路径转移。

OJ: luogu

题目 ID: P8820

难度:提高+/省选-

标签:图论最短路树形结构倍增矩阵

日期: 2026-07-06 08:46

题意

给定一棵 n 个点的树。两台主机之间的树上距离不超过 k 时,可以直接传输数据。一次请求从 st,可以选择若干中转主机,要求相邻中转主机之间都能直接传输。

经过的每台主机都要付出点权 v_i,包括起点和终点。要求每次询问的最小总代价。

思路

小数据可以先求出任意两点树上距离,再把所有距离不超过 k 的点对连边,最后在这个新图上求点权最短路:

cpp
// 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,而一次询问只关心树上 st 的路径。最优传输路线可以理解为:沿着这条路径从两端往中间走,允许一次跨过不超过 k 条边。

于是我们维护一个长度为 k 的 DP 状态:

text
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_stateright_state。最后枚举两侧末端距离 i,j,只要 i+j <= k,两边就可以接起来,取最小代价。

代码

cpp
// 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;
}

复杂度

预处理重链剖分和线段树为 O(nk3)O(n * k^3),这里 k <= 3,可以看成 O(n)O(n)

每次询问会跳过若干条重链,每次做常数大小矩阵或 DP 转移,复杂度为 O(log2nk3)O(log^2 n * k^3),可视为 O(log2n)O(log^2 n)

空间复杂度为 O(nk2)O(n * k^2)

总结

本题最基础的模型是“树的 k 次幂图上的点权最短路”。满分做法的关键,是发现 k <= 3 让“走过一段路径”可以压成很小的 min-plus 矩阵。

重链剖分负责把任意树上路径拆成少量连续链段,矩阵乘法负责把每段链上的选择/跳过决策合并起来。