[ZJOI2008] 树的统计

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

用树链剖分把树上路径拆成若干个 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 直接沿父亲链往上跳,直到两点相遇。 这个方法完全正确,但如果树退化成链,一次查询可能要走 O(n)O(n) 个点。

这题的核心是把树上路径转成区间问题。

树链剖分后:

  • 每个点有一个 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;
}

复杂度

预处理和建树是 O(n)O(n)

单点修改是 O(logn)O(log n)

每次路径查询会被拆成 O(logn)O(log n) 段,所以 QSUMQMAX 的复杂度都是 O(log2n)O(log^2 n)

空间复杂度是 O(n)O(n)

总结

这题是很标准的树链剖分模板题。

关键是两步:

  • 先把树上路径拆成若干个 DFS 序连续区间
  • 再在线段树上维护这些区间的统计信息

而且要记住负权这个细节,否则 QMAX 很容易写错。