[NOIP 2015 提高组] 运输计划

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

二分最长运输时间,树上差分检查所有超标路径是否共用一条足够长的边。

OJ: luogu

题目 ID: P2680

难度:提高

标签:二分答案LCA树上差分python

日期: 2026-07-17 02:00

题意

只能把一条边改为耗时 0,求所有运输计划同时完成的最短时间。

思路

二分答案 limit。对长度超过 limit 的路径做边差分,若一条边被所有超标路径共同经过,就可以同时缩短它们;该边权还必须至少覆盖最长缺口。逆序汇总差分即可找到所有公共边。

Python 知识

  • 运输路径先保存端点和长度,check 中反复使用同一批数据。
  • array("q") 保存距离和计划长度,array("i") 保存计数。
  • 二分模板只保留 feasible(mid) 一个判定函数。

代码

python
import sys
from array import array


input = sys.stdin.buffer.readline
n, plans = map(int, input().split())
head = array("i", [-1]) * (n + 1)
to = array("i", [0]) * (2 * n - 2)
next_edge = array("i", [0]) * (2 * n - 2)
edge_weight = array("i", [0]) * (2 * n - 2)
edge_count = 0


def add_edge(u, v, weight):
    global edge_count
    to[edge_count] = v
    edge_weight[edge_count] = weight
    next_edge[edge_count] = head[u]
    head[u] = edge_count
    edge_count += 1


for _ in range(n - 1):
    u, v, weight = map(int, input().split())
    add_edge(u, v, weight)
    add_edge(v, u, weight)

parent = array("i", [0]) * (n + 1)
depth = array("i", [0]) * (n + 1)
parent_weight = array("i", [0]) * (n + 1)
distance = array("q", [0]) * (n + 1)
order = array("i", [1])
index = 0
while index < n:
    node = order[index]
    index += 1
    edge = head[node]
    while edge != -1:
        neighbor = to[edge]
        if neighbor != parent[node]:
            parent[neighbor] = node
            depth[neighbor] = depth[node] + 1
            parent_weight[neighbor] = edge_weight[edge]
            distance[neighbor] = distance[node] + edge_weight[edge]
            order.append(neighbor)
        edge = next_edge[edge]

ancestors = [parent]
for _ in range(1, n.bit_length()):
    previous = ancestors[-1]
    ancestors.append(array("i", (previous[previous[node]] for node in range(n + 1))))


def lca(x, y):
    if depth[x] < depth[y]:
        x, y = y, x
    difference = depth[x] - depth[y]
    bit = 0
    while difference:
        if difference & 1:
            x = ancestors[bit][x]
        difference >>= 1
        bit += 1
    if x == y:
        return x
    for level in range(len(ancestors) - 1, -1, -1):
        if ancestors[level][x] != ancestors[level][y]:
            x = ancestors[level][x]
            y = ancestors[level][y]
    return parent[x]


u_values = array("i")
v_values = array("i")
lengths = array("q")
right_bound = 0
for _ in range(plans):
    u, v = map(int, input().split())
    ancestor = lca(u, v)
    length = distance[u] + distance[v] - 2 * distance[ancestor]
    u_values.append(u)
    v_values.append(v)
    lengths.append(length)
    right_bound = max(right_bound, length)

diff = array("i", [0]) * (n + 1)


def feasible(limit):
    for i in range(1, n + 1):
        diff[i] = 0
    bad = 0
    need = 0
    for i, length in enumerate(lengths):
        if length > limit:
            bad += 1
            need = max(need, length - limit)
            u, v = u_values[i], v_values[i]
            ancestor = lca(u, v)
            diff[u] += 1
            diff[v] += 1
            diff[ancestor] -= 2
    if not bad:
        return True
    best_edge = 0
    for node in reversed(order[1:]):
        if diff[node] == bad:
            best_edge = max(best_edge, parent_weight[node])
        diff[parent[node]] += diff[node]
    return best_edge >= need


left, right = 0, right_bound
while left < right:
    middle = (left + right) // 2
    if feasible(middle):
        right = middle
    else:
        left = middle + 1
print(left)

原有 C++ 版本仍保留:

cpp
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 300005;
const int LOG = 20;

struct Query {
    int u, v, lca_node;
    long long length;
};

int n, m;
int head[MAXN], to[MAXN * 2], nxt[MAXN * 2], edge_weight[MAXN * 2], edge_cnt;
int depth_node[MAXN];
int up[MAXN][LOG + 1];          // up[x][j] 表示 x 的 2^j 级祖先。
int parent_edge_weight[MAXN];   // parent_edge_weight[x] 表示 x 到父亲这条边的权值。
long long dist_root[MAXN];      // 根到每个点的路径长度。
int diff_count[MAXN];           // check() 中的边差分计数。
vector<int> bfs_order;          // BFS 顺序,反向使用即可自底向上汇总。
Query query_data[MAXN];

void add_edge(int u, int v, int w) {
    edge_cnt++;
    to[edge_cnt] = v;
    edge_weight[edge_cnt] = w;
    nxt[edge_cnt] = head[u];
    head[u] = edge_cnt;
}

void read_tree() {
    cin >> n >> m;
    for (int i = 1; i < n; i++) {
        int u, v, w;
        cin >> u >> v >> w;
        add_edge(u, v, w);
        add_edge(v, u, w);
    }
}

void build_lca() {
    queue<int> que;
    que.push(1);
    depth_node[1] = 1;
    bfs_order.reserve(n);

    // 用 BFS 建父亲、深度和根到点距离,避免深递归爆栈。
    while (!que.empty()) {
        int u = que.front();
        que.pop();
        bfs_order.push_back(u);

        for (int j = 1; j <= LOG; j++) {
            up[u][j] = up[up[u][j - 1]][j - 1];
        }

        for (int i = head[u]; i != 0; i = nxt[i]) {
            int v = to[i];
            if (v == up[u][0]) {
                continue;
            }
            up[v][0] = u;
            depth_node[v] = depth_node[u] + 1;
            dist_root[v] = dist_root[u] + edge_weight[i];
            parent_edge_weight[v] = edge_weight[i];
            que.push(v);
        }
    }
}

int lca(int x, int y) {
    if (depth_node[x] < depth_node[y]) {
        swap(x, y);
    }

    int diff = depth_node[x] - depth_node[y];
    for (int j = LOG; j >= 0; j--) {
        if ((diff & (1 << j)) != 0) {
            x = up[x][j];
        }
    }

    if (x == y) {
        return x;
    }

    for (int j = LOG; j >= 0; j--) {
        if (up[x][j] != up[y][j]) {
            x = up[x][j];
            y = up[y][j];
        }
    }

    return up[x][0];
}

void read_queries_and_precompute(long long &right_bound) {
    right_bound = 0;
    for (int i = 1; i <= m; i++) {
        int u, v;
        cin >> u >> v;
        int g = lca(u, v);
        long long len = dist_root[u] + dist_root[v] - 2LL * dist_root[g];
        query_data[i] = {u, v, g, len};
        right_bound = max(right_bound, len);
    }
}

bool check(long long limit) {
    for (int i = 1; i <= n; i++) {
        diff_count[i] = 0;
    }

    int bad_count = 0;
    long long need_reduce = 0;

    for (int i = 1; i <= m; i++) {
        if (query_data[i].length <= limit) {
            continue;
        }

        bad_count++;
        need_reduce = max(need_reduce, query_data[i].length - limit);

        // 只统计超出 limit 的路径。若要一次改边让它们都变短,
        // 这条边必须被所有超标路径共同经过。
        int u = query_data[i].u;
        int v = query_data[i].v;
        int g = query_data[i].lca_node;
        diff_count[u]++;
        diff_count[v]++;
        diff_count[g] -= 2;
    }

    if (bad_count == 0) {
        return true;
    }

    long long best_common_edge = 0;

    // 反向 BFS 序等价于从叶子向根汇总。
    // diff_count[x] 汇总后表示边 parent[x] - x 被多少条超标路径经过。
    for (int i = (int)bfs_order.size() - 1; i >= 0; i--) {
        int u = bfs_order[i];
        if (diff_count[u] == bad_count) {
            best_common_edge = max(best_common_edge, (long long)parent_edge_weight[u]);
        }
        if (up[u][0] != 0) {
            diff_count[up[u][0]] += diff_count[u];
        }
    }

    // 最长的超标路径至少需要被缩短 need_reduce。
    return best_common_edge >= need_reduce;
}

void solve() {
    build_lca();

    long long right_bound = 0;
    read_queries_and_precompute(right_bound);

    long long left = 0;
    long long right = right_bound;
    long long answer = right_bound;

    while (left <= right) {
        long long mid = (left + right) / 2;
        if (check(mid)) {
            answer = mid;
            right = mid - 1;
        } else {
            left = mid + 1;
        }
    }

    cout << answer << '\n';
}

int main() {
    ios::sync_with_stdio(false);
    cin.tie(nullptr);

    read_tree();
    solve();

    return 0;
}

复杂度

预处理 O((n+m)log n),每次判定 O((n+m)log n),总复杂度 O((n+m)log n log W)

总结

“改一条边”意味着所有超标路径必须有公共边,树上差分正好能检测这个交集。