【P4427 [BJOI2018] 求和】 树上前缀和

一、题目模型简述

给定一棵无根树,共 $n$ 个节点,多组询问,每组给出三个参数 $u,v,k$:

求树上 $u$ 到 $v$ 这条路径上所有点的深度的 $k$ 次方之和,结果对 $998244353$ 取模。

核心前置知识

  1. 倍增LCA:快速求树上两点最近公共祖先;
  2. 树上前缀和:定义 $s[u][k]$ 为根节点到节点 $u$ 整条路径所有点深度的 $k$ 次幂总和
  3. 树上路径拆分公式:设 $l = \text{lca}(u,v)$,$f = fa[l][0]$($l$ 的父节点)

$$
\text{sum}(u \to v) = s[u][k] + s[v][k] - s[l][k] - s[f][k]
$$

原理

  • $s[u]$:根到 $u$;$s[v]$:根到 $v$;
  • 两段相加后,根~$l$ 的路径被重复计算2次,需要减去 $2\times s[l]$;
  • 等价于减去 $s[l]$(u侧重复段)和 $s[fa[l]]$(v侧重复段)。
    请添加图片描述

二、整体算法思路

1. 预处理阶段

  1. 邻接表存树;
  2. DFS遍历整棵树,同步完成两件预处理:
    • 倍增数组 $fa[N][LOG]$:存储每个节点 $2^i$ 级祖先,用于LCA;
    • 树上前缀和数组 $s[N][KMAX]$:$s[x][j]$ 存根到 $x$ 所有点深度的 $j$ 次幂和;
  3. 预处理每个节点深度的 $1\sim50$ 次幂,模 $998244353$ 存入临时数组,再累加进前缀和。

2. 查询阶段

每组询问执行三步:

  1. 调用倍增LCA求出 $u,v$ 的公共祖先 $l$;
  2. 套用树上路径求和公式计算总答案;
  3. 处理模运算负数问题,输出正数结果。

3. 数据范围说明

  • 节点总数 $n \le 3\times 10^5$;
  • 幂次 $1\le k \le 50$;
  • 倍增层数LOG=22,满足 $2^{20}>3\times 10^5$;
  • 模数固定 $mod=998244353$,全程使用long long防止乘法溢出。

三、完整AC代码

#include <iostream>
#include <cstdio>
#include <algorithm>
using namespace std;

typedef long long LL;
const int N = 300010;    // 节点上限3e5
const int LOG = 22;
const int KMAX = 50;
const int mod = 998244353;

// 倍增LCA数组 fa[u][i]:u的2^i倍祖先
int fa[N][LOG];
int dep[N];               // dep[u]:节点u的深度
LL mi[KMAX + 5];         // 临时数组,存当前点深度1~50次幂
LL s[N][KMAX + 5];       // s[u][k]:根到u路径所有点深度k次方前缀和

// 链式前向星邻接表
int h[N], to[N << 1], ne[N << 1];
int tot = 0;

// 加双向边
void add(int a, int b) {
    to[++tot] = b;
    ne[tot] = h[a];
    h[a] = tot;
}

// DFS:预处理倍增祖先、深度、树上幂次前缀和
void dfs(int u, int f) {
    // 递推倍增祖先
    for (int i = 1; i <= 20; i++) {
        fa[u][i] = fa[fa[u][i - 1]][i - 1];
    }
    
    // 遍历所有邻接点
    for (int i = h[u]; i; i = ne[i]) {
        int v = to[i];
        if (v == f) continue;
        
        fa[v][0] = u;
        dep[v] = dep[u] + 1;
        
        // 计算当前深度dep[v]的1~50次幂
        mi[0] = 1;
        for (int j = 1; j <= 50; j++) {
            mi[j] = mi[j - 1] * dep[v] % mod;
        }
        
        // 根到v = 根到u的和 + 当前点深度的k次幂
        for (int j = 1; j <= 50; j++) {
            s[v][j] = (mi[j] + s[u][j]) % mod;
        }
        
        dfs(v, u);
    }
}

// 倍增求u、v的最近公共祖先LCA
int lca(int u, int v) {
    // 保证u深度更深
    if (dep[u] < dep[v]) swap(u, v);
    
    // u向上跳到和v同一深度
    for (int i = 20; i >= 0; i--) {
        if (dep[fa[u][i]] >= dep[v]) {
            u = fa[u][i];
        }
    }
    
    if (u == v) return v;
    
    // 两点同步向上跳,直到LCA下一层
    for (int i = 20; i >= 0; i--) {
        if (fa[u][i] != fa[v][i]) {
            u = fa[u][i];
            v = fa[v][i];
        }
    }
    return fa[u][0];
}

int main() {
    int n;
    scanf("%d", &n);
    
    // 读入n-1条树边,双向建图
    for (int i = 1; i <= n - 1; i++) {
        int a, b;
        scanf("%d%d", &a, &b);
        add(a, b);
        add(b, a);
    }
    
    // 根节点1初始化,根深度为0,0的任意次幂为0
    dep[1] = 0;
    for (int j = 1; j <= 50; j++) {
        s[1][j] = 0;
    }
    fa[1][0] = 0; // 根没有父节点
    
    dfs(1, 0);
    
    int m;
    scanf("%d", &m);
    
    // 处理m组询问
    while (m--) {
        int u, v, k;
        scanf("%d%d%d", &u, &v, &k);
        
        int l = lca(u, v);
        
        // 树上路径求和公式
        LL ans = (s[u][k] + s[v][k] - s[l][k] - s[fa[l][0]][k]) % mod;
        
        // 减法会出现负数,加两倍模数保证结果为正再取模
        ans = (ans + 2 * mod) % mod;
        printf("%lld\n", ans);
    }
    
    return 0;
}

四、代码分段详细讲解

1. 常量与全局数组定义

const int N = 300010;
const int LOG = 22;
const int KMAX = 50;
const int mod = 998244353;

int fa[N][LOG];
int dep[N];
LL mi[KMAX + 5];
LL s[N][KMAX + 5];
int h[N], to[N << 1], ne[N << 1];
int tot = 0;
  • N=3e5:适配题目节点上限;LOG=22 倍增层数,覆盖最大深度;
  • KMAX=50:题目幂次最大为50,预处理1~50次幂;
  • fa[N][LOG]:倍增祖先数组,fa[u][i] 代表u向上跳 $2^i$ 步到达的节点;
  • dep[]:存储每个节点的深度,根节点1深度定义为0;
  • mi[]:临时数组,单次DFS内计算当前点深度各次幂;
  • s[N][KMAX]:核心树上前缀和数组,s[x][k] = 根1到x路径上所有点深度的k次方之和;
  • 邻接表数组 h,to,neN<<1 开双倍空间,存储双向树边。

2. 邻接表加边函数 add

void add(int a, int b) {
    to[++tot] = b;
    ne[tot] = h[a];
    h[a] = tot;
}

树是无向图,输入每条边需要执行两次add,分别建立双向连通关系。

3. DFS预处理函数

void dfs(int u, int f) {
    // 倍增祖先递推
    for (int i = 1; i <= 20; i++) {
        fa[u][i] = fa[fa[u][i - 1]][i - 1];
    }
    
    // 遍历子节点
    for (int i = h[u]; i; i = ne[i]) {
        int v = to[i];
        if (v == f) continue;
        
        fa[v][0] = u;
        dep[v] = dep[u] + 1;
        
        // 计算深度1~50次幂
        mi[0] = 1;
        for (int j = 1; j <= 50; j++) {
            mi[j] = mi[j - 1] * dep[v] % mod;
        }
        
        // 前缀和传递:根到v = 根到u + 当前点贡献
        for (int j = 1; j <= 50; j++) {
            s[v][j] = (mi[j] + s[u][j]) % mod;
        }
        
        dfs(v, u);
    }
}
  1. 倍增递推:已知 $2^{i-1}$ 级祖先,推出 $2^i$ 级祖先;
  2. 跳过父节点避免回溯;子节点父节点设为u,深度+1;
  3. 幂次计算:循环算出当前深度的1~50次幂,全程取模防止溢出;
  4. 树上前缀和:从父节点的前缀和累加当前点幂次,得到根到当前节点的前缀和。

4. 倍增LCA函数

int lca(int u, int v) {
    // 保证u深度更深
    if (dep[u] < dep[v]) swap(u, v);
    
    // u向上跳到和v同一深度
    for (int i = 20; i >= 0; i--) {
        if (dep[fa[u][i]] >= dep[v]) {
            u = fa[u][i];
        }
    }
    
    if (u == v) return v;
    
    // 两点同步向上跳,直到LCA下一层
    for (int i = 20; i >= 0; i--) {
        if (fa[u][i] != fa[v][i]) {
            u = fa[u][i];
            v = fa[v][i];
        }
    }
    return fa[u][0];
}
  1. 深度对齐:让深度大的节点u向上跳,直到与v同深度;
  2. 同步上跳:如果u和v不同,则同时向上跳 $2^i$ 步,直到它们的父节点相同;
  3. 返回LCA:最终返回u的父节点即为最近公共祖先。

5. 主函数与查询处理

int main() {
    // ... 建图、DFS预处理 ...
    
    int m;
    scanf("%d", &m);
    
    // 处理m组询问
    while (m--) {
        int u, v, k;
        scanf("%d%d%d", &u, &v, &k);
        
        int l = lca(u, v);
        
        // 树上路径求和公式
        LL ans = (s[u][k] + s[v][k] - s[l][k] - s[fa[l][0]][k]) % mod;
        
        // 减法会出现负数,加两倍模数保证结果为正再取模
        ans = (ans + 2 * mod) % mod;
        printf("%lld\n", ans);
    }
    
    return 0;
}

核心公式解析

  • s[u][k]:根到u路径上所有点深度的k次方和
  • s[v][k]:根到v路径上所有点深度的k次方和
  • s[l][k]:根到LCA路径上所有点深度的k次方和
  • s[fa[l][0]][k]:根到LCA父节点路径上所有点深度的k次方和

公式推导

路径u→v = (根→u) + (根→v) - 2×(根→l) + l
        = s[u] + s[v] - 2×s[l] + dep[l]^k
        = s[u] + s[v] - s[l] - (s[l] - dep[l]^k)
        = s[u] + s[v] - s[l] - s[fa[l]]

五、时间复杂度分析

  1. 预处理阶段

    • DFS遍历:$O(n)$
    • 每个节点计算50次幂:$O(50n)$
    • 倍增数组预处理:$O(n \log n)$
    • 总复杂度:$O(n \log n + 50n)$
  2. 查询阶段

    • 每次LCA查询:$O(\log n)$
    • 公式计算:$O(1)$
    • m次查询总复杂度:$O(m \log n)$
  3. 空间复杂度

    • 邻接表:$O(n)$
    • 倍增数组:$O(n \log n)$
    • 前缀和数组:$O(50n)$

六、常见问题与优化

Q1:为什么根节点深度设为0?

  • 方便计算:深度从0开始,0的任意次幂为0,不影响前缀和结果;
  • 避免特殊处理:如果深度从1开始,根节点的深度幂次需要单独处理。

Q2:为什么公式中要减去s[fa[l][0]]而不是s[l]?

  • 因为 $s[l]$ 包含了LCA节点本身的贡献,而路径u→v只需要计算一次LCA节点的深度幂次;
  • $s[fa[l][0]]$ 是根到LCA父节点的前缀和,不包含LCA节点;
  • 所以减去 $s[l]$ 和 $s[fa[l][0]]$ 相当于减去两次根到LCA父节点的贡献,再加上一次LCA节点的贡献。

Q3:如何处理模运算的负数?

  • C++中 % 运算符对负数取模结果仍为负数;
  • 使用 (ans + 2 * mod) % mod 确保结果在 $[0, mod-1]$ 范围内。

Q4:能否进一步优化?

  • 对于固定k值,可以预处理深度幂次表,避免DFS中重复计算;
  • 使用更高效的LCA算法(如树链剖分)可略微提升查询速度;
  • 对于特别大的n,可考虑使用内存更紧凑的数据结构。

七、总结

本题是树上路径统计的经典问题,结合了:

  1. 倍增LCA:快速求最近公共祖先
  2. 树上前缀和:预处理根到每个节点的路径信息
  3. 路径拆分公式:将任意路径转化为根到两点的路径差

关键技巧

  • 预处理所有可能的k值(1~50)的前缀和,避免每次查询重新计算;
  • 使用模运算处理大数,注意负数修正;
  • 深度从0开始定义,简化边界处理。

掌握这种预处理+路径拆分的思想,可以解决许多类似的树上路径统计问题。

posted on 2026-07-12 16:42  5iCode  阅读(4)  评论(0)    收藏  举报

导航