[JLOI2014] 松鼠的新家

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

对参观顺序中的相邻点路径做树上点差分,汇总后减去每段交界点的重复计数。

OJ: luogu

题目 ID: P3258

难度:普及+/提高

标签:树上差分LCA倍增

日期: 2026-06-22 22:39

题意

给定一棵树和参观顺序 a_1..a_n。维尼依次从 a_i 走到 a_{i+1},每走到一个房间就吃一块糖。

最后到达 a_n 时不吃糖。要求每个房间至少放多少糖。

思路

先看一个可以直接验证想法的朴素解:

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

// brute.cpp:逐条路径 BFS 找父亲并枚举路径点,只适合小数据。

const int MAXN = 505;

int n;
int route_node[MAXN];
vector<int> graph_edges[MAXN];
long long answer[MAXN];
int parent_node[MAXN];

void mark_path(int start, int target) {
    for (int i = 1; i <= n; i++) {
        parent_node[i] = -1;
    }
    queue<int> que;
    que.push(start);
    parent_node[start] = 0;

    while (!que.empty()) {
        int u = que.front();
        que.pop();
        if (u == target) {
            break;
        }
        for (int i = 0; i < (int)graph_edges[u].size(); i++) {
            int v = graph_edges[u][i];
            if (parent_node[v] == -1) {
                parent_node[v] = u;
                que.push(v);
            }
        }
    }

    int x = target;
    while (x != 0) {
        answer[x]++;
        if (x == start) {
            break;
        }
        x = parent_node[x];
    }
}

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

    cin >> n;
    for (int i = 1; i <= n; i++) {
        cin >> route_node[i];
    }
    for (int i = 1; i < n; i++) {
        int u, v;
        cin >> u >> v;
        graph_edges[u].push_back(v);
        graph_edges[v].push_back(u);
    }

    for (int i = 1; i < n; i++) {
        mark_path(route_node[i], route_node[i + 1]);
    }
    for (int i = 2; i <= n; i++) {
        answer[route_node[i]]--;
    }

    for (int i = 1; i <= n; i++) {
        cout << answer[i] << '\n';
    }

    return 0;
}

暴力逐条路径枚举经过的点会超时。这里是典型的树上路径点加一,可以用树上点差分。

对一条路径 u -> v,设 g = lca(u, v)。点差分标记为:

text
diff[u]++
diff[v]++
diff[g]--
diff[parent[g]]--

最后自底向上把子树差分累加起来,每个点得到被多少条路径经过。

还要处理一个细节:a_i 是上一段路径的终点,也是下一段路径的起点,路径差分会把它算两次,但实际只到达一次。因此最后对 a_2..a_n 各减一。这里包括 a_n,正好对应最后到餐厅不吃糖。

代码

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

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

int n;
int route_node[MAXN];
int head[MAXN], to[MAXN * 2], nxt[MAXN * 2], edge_cnt;
int depth_node[MAXN];
int up[MAXN][LOG + 1];
long long diff_count[MAXN];
long long answer[MAXN];

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

void read_input() {
    cin >> n;
    for (int i = 1; i <= n; i++) {
        cin >> route_node[i];
    }
    for (int i = 1; i < n; i++) {
        int u, v;
        cin >> u >> v;
        add_edge(u, v);
        add_edge(v, u);
    }
}

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

    while (!que.empty()) {
        int u = que.front();
        que.pop();

        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;
            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 collect_answer() {
    vector<int> order;
    order.reserve(n);
    queue<int> que;
    que.push(1);
    while (!que.empty()) {
        int u = que.front();
        que.pop();
        order.push_back(u);
        for (int i = head[u]; i != 0; i = nxt[i]) {
            int v = to[i];
            if (v == up[u][0]) {
                continue;
            }
            que.push(v);
        }
    }

    for (int i = (int)order.size() - 1; i >= 0; i--) {
        int u = order[i];
        answer[u] += diff_count[u];
        if (up[u][0] != 0) {
            diff_count[up[u][0]] += diff_count[u];
        }
    }
}

void solve() {
    build_lca();

    for (int i = 1; i < n; i++) {
        int u = route_node[i];
        int v = route_node[i + 1];
        int g = lca(u, v);
        diff_count[u]++;
        diff_count[v]++;
        diff_count[g]--;
        if (up[g][0] != 0) {
            diff_count[up[g][0]]--;
        }
    }

    collect_answer();

    // 每个中间到达点既是上一段终点又是下一段起点,只应拿一次糖。
    for (int i = 2; i <= n; i++) {
        answer[route_node[i]]--;
    }

    for (int i = 1; i <= n; i++) {
        cout << answer[i] << '\n';
    }
}

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

    read_input();
    solve();

    return 0;
}

复杂度

LCA 预处理和路径处理总时间复杂度为 O(nlogn)O(n log n)

空间复杂度为 O(nlogn)O(n log n)

总结

多次树上路径点加一,优先考虑树上点差分。

本题最容易漏的是相邻路径交界点的重复计数,最后必须对 a_2..a_n 做一次修正。