[BJOI2018] 求和

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

对每个节点预处理根到它路径上的 depth^k 前缀和,再用 LCA 把路径拆成两段:sum(x)+sum(y)-2sum(lca)+depth(lca)^k。

OJ: luogu

题目 ID: P4427

难度:提高+/省选-

标签:LCA倍增树形结构

日期: 2026-06-20 02:44

题意

给一棵以 1 为根的树。

每次询问给出 x, y, k,要求计算:

  • xy 这条路径上
  • 所有节点深度的 k 次方和

其中深度定义为:

  • 节点到根 1 的路径边数

结果对 998244353 取模。

样例树

样例树结构如下:

graph G {
  1 -- 2;
  1 -- 3;
  2 -- 4;
  2 -- 5;
}

深度分别是:

  • dep[1]=0dep[1]=0
  • dep[2]=1dep[2]=1
  • dep[3]=1dep[3]=1
  • dep[4]=2dep[4]=2
  • dep[5]=2dep[5]=2

比如查询 4, 5, 1,路径是 4-2-5,答案就是:

  • 21+11+21=52^1 + 1^1 + 2^1 = 5

思路

先看一个最直接的小数据暴力:

cpp
// brute.cpp:每次查询直接找出 x 到 y 的唯一路径,然后把路径上所有点的 depth^k 加起来。
// 这个做法复杂度高,只适合小数据理解和对拍。
#include <bits/stdc++.h>
using namespace std;

const int MAXN = 25;
const int MOD = 998244353;

int n, m;
vector<int> g[MAXN];
int depth_arr[MAXN];
int parent_arr[MAXN];

void build_depth() {
    vector<int> st;
    st.push_back(1);
    parent_arr[1] = 0;
    depth_arr[1] = 0;

    while (!st.empty()) {
        int u = st.back();
        st.pop_back();

        for (size_t i = 0; i < g[u].size(); i++) {
            int v = g[u][i];
            if (v == parent_arr[u]) {
                continue;
            }
            parent_arr[v] = u;
            depth_arr[v] = depth_arr[u] + 1;
            st.push_back(v);
        }
    }
}

bool dfs_find(int u, int target, int fa) {
    if (u == target) {
        return true;
    }

    for (size_t i = 0; i < g[u].size(); i++) {
        int v = g[u][i];
        if (v == fa) {
            continue;
        }
        parent_arr[v] = u;
        if (dfs_find(v, target, u)) {
            return true;
        }
    }
    return false;
}

int power_mod(int a, int k) {
    long long res = 1;
    for (int i = 1; i <= k; i++) {
        res = res * a % MOD;
    }
    return (int)res;
}

int query_path_sum(int x, int y, int k) {
    for (int i = 1; i <= n; i++) {
        parent_arr[i] = 0;
    }
    dfs_find(x, y, 0);

    int ans = 0;
    int u = y;
    while (u != x) {
        ans += power_mod(depth_arr[u], k);
        if (ans >= MOD) {
            ans -= MOD;
        }
        u = parent_arr[u];
    }
    ans += power_mod(depth_arr[x], k);
    if (ans >= MOD) {
        ans -= MOD;
    }
    return ans;
}

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

    cin >> n;
    for (int i = 1; i <= n; i++) {
        g[i].clear();
        depth_arr[i] = 0;
        parent_arr[i] = 0;
    }

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

    build_depth();

    cin >> m;
    while (m--) {
        int x, y, k;
        cin >> x >> y >> k;
        cout << query_path_sum(x, y, k) << '\n';
    }

    return 0;
}

暴力做法就是:

  1. 每次查询找出 x -> y 的唯一路径
  2. 枚举路径上的每个点
  3. 把它的 depthkdepth^k 累加起来

这个方法最贴近题意,但查询多时会很慢。

这题的关键是把“路径和”拆成“根到点前缀和”。

设:

  • sumpow[u][k]sum_pow[u][k] 表示从根到 u 的路径上,所有节点 depthkdepth^k 的和

那么对查询 (x, y, k),设 p = lca(x, y),就有:

  • 根到 x 的前缀和:sumpow[x][k]sum_pow[x][k]
  • 根到 y 的前缀和:sumpow[y][k]sum_pow[y][k]
  • 根到 p 这一段被重复算了两次

所以答案自然是:

sumpow[x][k]+sumpow[y][k]2sumpow[p][k]+depth[p]ksum_pow[x][k] + sum_pow[y][k] - 2 * sum_pow[p][k] + depth[p]^k

最后为什么还要加回 depth[p]kdepth[p]^k

因为:

  • p 本身属于真实路径
  • 但它在前面的减法里被减掉了两次

于是只要预处理好两样东西:

  1. 倍增 LCA
  2. 对每个节点、每个 k(1..50)k (1..50) 的前缀和

每次查询就能很快回答。

由于 k 的范围只有 50,所以对每个节点把 depth1,depth2,,depth50depth^1, depth^2, \dots, depth^{50} 全部顺手算出来是完全可行的。

代码

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

const int MAXN = 300000 + 5;
const int MAXM = 600000 + 5;
const int LOG = 20;
const int MAXK = 50;
const int MOD = 998244353;

int n, m;
int head[MAXN], to[MAXM], nxt[MAXM], edge_cnt;

int depth_arr[MAXN];
int up[MAXN][LOG];
int sum_pow[MAXN][MAXK + 1];  // 根到当前点路径上,depth^k 的前缀和

void init_graph(int n) {
    edge_cnt = 0;
    for (int i = 1; i <= n; i++) {
        head[i] = 0;
        depth_arr[i] = 0;
        for (int j = 0; j < LOG; j++) {
            up[i][j] = 0;
        }
        for (int k = 1; k <= MAXK; k++) {
            sum_pow[i][k] = 0;
        }
    }
}

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

void build_lca_and_prefix(int root) {
    vector<int> st;
    st.push_back(root);

    while (!st.empty()) {
        int u = st.back();
        st.pop_back();

        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_arr[v] = depth_arr[u] + 1;

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

            long long pw = 1;
            for (int k = 1; k <= MAXK; k++) {
                pw = pw * depth_arr[v] % MOD;
                sum_pow[v][k] = sum_pow[u][k] + (int)pw;
                if (sum_pow[v][k] >= MOD) {
                    sum_pow[v][k] -= MOD;
                }
            }

            st.push_back(v);
        }
    }
}

int kth_ancestor(int u, int k) {
    for (int j = 0; j < LOG; j++) {
        if (k & (1 << j)) {
            u = up[u][j];
        }
    }
    return u;
}

int lca(int a, int b) {
    if (depth_arr[a] < depth_arr[b]) {
        swap(a, b);
    }

    a = kth_ancestor(a, depth_arr[a] - depth_arr[b]);
    if (a == b) {
        return a;
    }

    for (int j = LOG - 1; j >= 0; j--) {
        if (up[a][j] != up[b][j]) {
            a = up[a][j];
            b = up[b][j];
        }
    }

    return up[a][0];
}

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

    cin >> n;
    init_graph(n);

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

    build_lca_and_prefix(1);

    cin >> m;
    while (m--) {
        int x, y, k;
        cin >> x >> y >> k;

        int p = lca(x, y);

        long long depth_pow = 1;
        for (int i = 1; i <= k; i++) {
            depth_pow = depth_pow * depth_arr[p] % MOD;
        }

        int ans = sum_pow[x][k];
        ans += sum_pow[y][k];
        if (ans >= MOD) {
            ans -= MOD;
        }

        ans -= 2LL * sum_pow[p][k] % MOD;
        if (ans < 0) {
            ans += MOD;
        }
        ans += depth_pow;
        if (ans >= MOD) {
            ans -= MOD;
        }

        cout << ans << '\n';
    }

    return 0;
}

复杂度

预处理:

  • 倍增祖先表:O(nlogn)O(n log n)
  • 对每个点计算 1..50 次幂前缀和:O(50n)O(50n)

每次查询:

  • 求一次 LCA:O(logn)O(log n)
  • 再算一次 depth[lca]kdepth[lca]^kO(k)O(k),这里 k50k \leqslant 50

总复杂度可以写成:

  • O(nlogn+50n+m(logn+50))O(n \log n + 50n + m(\log n + 50))

空间复杂度:

  • O(nlogn+50n)O(n log n + 50n)

总结

这题最重要的观察是:

  • 虽然查询里的 k 会变化,但它只在 1..50 之间

这意味着我们完全可以把每个点对应的 depthkdepth^k 全部预处理掉。

于是整题就变成非常标准的套路:

  1. 倍增求 LCA
  2. 根到点前缀和
  3. sum(x)+sum(y)2sum(lca)+val(lca)sum(x)+sum(y)-2sum(lca)+val(lca) 还原路径和

本质上是一道:

  • LCA + 树上前缀和

的组合题。

一图流解析

这张图把本题的建模、关键转移、实现检查和训练方法压缩到一页,适合读完正文后复盘。

一图流解析