题解:P16554 [ICPC 2026 LAC] Holes and Tunnels
lailai0916
·
·
题解
题意简述
给定一棵 n 个节点的树。路线是两个不同节点之间的无向简单路径,两名玩家各选择一条路线。
对每个 1\le k<n,统计两条路线恰有 k 条公共边的有序路线对数,对 998244353 取模。
解题思路
两条树上简单路径的公共边若非空,必然构成一条简单路径。因为公共部分中任意两个节点之间的唯一路径,同时属于原来的两条路径,不可能断开。
因此,可以先固定公共路径的两个不同端点 u,v,计算恰好以这条路径为公共部分的路线对数,再按 \operatorname{dist}(u,v) 累加。
考虑公共路径在端点 u 的延伸情况。设 z 是从 u 走向 v 的第一个邻居。删去节点 u 后,每个邻居 x 对应一个连通块,记其大小为 s_{u,x}。
每名玩家的路线在 u 这一侧,都要选择一个端点:可以停在 u,也可以进入除 z 以外的任意分支。这一侧共有 n-s_{u,z} 个可选节点。两名玩家有顺序,因此最初有 (n-s_{u,z})^2 种端点选择。
若两人都选中同一个分支 x\ne z 内的节点,两条路线就会共同经过边 ux,公共路径会继续向外延伸,必须排除。各分支对应的排除情况互不相交;两人选择不同分支,或至少一人停在 u,都不会增加公共边。因此,端点 u 的合法选择数为:
W(u,z)=(n-s_{u,z})^2-\sum_{x\ne z}s_{u,x}^2
端点 v 也有同样的计算。公共路径两端可供延伸的区域互不相交,选择相互独立,所以固定公共路径的贡献就是两端的 W 之积。
每个这样的端点选择确实确定了两条合法路线:每名玩家在 u 侧和 v 侧各有一个端点,它们之间的唯一路径必然包含 u 到 v 的公共路径。两端的排除条件又保证不会多出公共边。反过来,任何具有该公共路径的路线对也都能恢复出这四个端点。因此,乘积既没有漏计,也没有重复计算路线对。
为了快速计算每个方向的 W,记:
A_u=n^2-\sum_x s_{u,x}^2
对于边 uz,删边后两侧节点数为 s 和 n-s,再记 B_{uz}=2s(n-s)。将上面的式子展开,得到:
\begin{aligned}
W(u,z) & =n^2-2ns_{u,z}+s_{u,z}^2-\sum_x s_{u,x}^2+s_{u,z}^2 \\
& =A_u-2s_{u,z}(n-s_{u,z}) \\
& =A_u-B_{uz}
\end{aligned}
剩下的问题是:枚举无序节点对 $\set{u,v}$,将两端相应方向的权值乘积加入距离对应的答案。使用点分治,并在每个重心处进行多项式卷积。
设当前重心为 $c$,删去它后得到若干连通块。对于其中一个连通块 $T_i$,定义多项式:
$$
H_i(x)=\sum_{u\in T_i}W(u,p_u)x^{\operatorname{dist}(c,u)}
$$
其中 $p_u$ 是从 $u$ 走向 $c$ 的第一个邻居。固定重心和所在连通块后,这个方向就已确定,与公共路径另一端的具体位置无关。
若 $u,v$ 分属两个不同连通块,则它们之间的路径经过 $c$,且距离等于两者到 $c$ 的距离之和。因此,卷积能按公共边数汇总这类贡献。令 $H=\sum_i H_i$,则:
$$
H(x)^2-\sum_i H_i(x)^2
$$
恰好保留两端位于不同连通块的项。平方中的端点有顺序,所以无序端点对 $\set{u,v}$ 会出现两次。
还需要计算公共路径的一端就是重心的情况。对于连通块 $T_i$,设它与 $c$ 相邻的节点为 $z_i$,则这类贡献为 $W(c,z_i)H_i(x)$。重心的权值依赖延伸方向,不能给所有连通块统一乘上一个点权。
为了让两类贡献使用相同倍数,代码把当前重心的贡献写为:
$$
H(x)^2-\sum_i H_i(x)^2+2\sum_i W(c,z_i)H_i(x)
$$
完成全部点分治后,再将每个系数乘以 $2$ 的模逆元。这里消除的是公共路径两个端点的重复枚举,**不是**交换两名玩家产生的重复;玩家的顺序已经包含在两端 $W$ 的计数中,必须保留。
对于两个端点,点分治中第一次把它们分入不同连通块,或选中其中一个端点作为重心时,就会且仅会统计这条公共路径。仍处于同一连通块的端点对在本层被减去,交给后续递归处理,所以全部公共路径恰好统计一次。
特别注意,所有 $A_u$ 和 $B_{uz}$ 都必须始终使用原树的大小。点分治划分的是公共路径的端点范围,并没有禁止玩家的完整路线向已删除重心所在的区域延伸。若重新按当前连通块大小计算 $W$,就会漏掉这些合法路线对。
代码中,`val` 保存 $A$,邻接表第二个字段保存 $B$。`solve` 先求当前连通块的重心,然后依次遍历每个分支,用 `dep` 统计距离,用 `w` 记录走向重心的第一条边权。数组 `b` 保存当前 $H_i$,数组 `a` 累加得到 $H$;`square` 将多项式平方按指定正负号加入答案。
多项式较短时直接相乘,较长时使用快速数论变换(Number Theoretic Transform,NTT)。模数 $998244353$ 支持所需的变换长度,取原根 $3$。所有按树边遍历的过程都使用数组队列,递归仅用于点分治,避免链状树产生线性深度的递归。
设当前连通块大小为 $s$。各分支多项式次数之和不超过 $s-1$,总多项式次数也不超过 $s-1$,所以本层卷积需要 $O(s\log s)$ 时间。点分治每下降一层,连通块大小至少减半,故总时间复杂度为 $O(n\log^2 n)$,空间复杂度为 $O(n)$。
## 参考代码
```cpp
#include <bits/stdc++.h>
using namespace std;
using ll=long long;
const int N=200005;
const int M=524293;
const int mod=998244353;
int n;
int fa[N],siz[N],q[N],dep[N],w[N],val[N],ans[N],a[N],b[N],f[M];
bool vis[N];
vector<pair<int,int>> G[N];
ll Pow(ll x,ll y)
{
x%=mod;
ll res=1;
while(y)
{
if(y&1)res=res*x%mod;
x=x*x%mod;
y>>=1;
}
return res;
}
void ntt(int n,bool op)
{
for(int i=1,j=0;i<n;i++)
{
int k=n>>1;
while(j>=k){j^=k;k>>=1;}
j^=k;
if(i<j)swap(f[i],f[j]);
}
for(int i=2;i<=n;i<<=1)
{
int w=Pow(3,(mod-1)/i);
if(op)w=Pow(w,mod-2);
for(int j=0;j<n;j+=i)
{
ll cur=1;
for(int k=j;k<j+i/2;k++)
{
int x=f[k],y=cur*f[k+i/2]%mod;
f[k]=(x+y)%mod;
f[k+i/2]=(x-y+mod)%mod;
cur=cur*w%mod;
}
}
}
if(op)
{
int inv=Pow(n,mod-2);
for(int i=0;i<n;i++)f[i]=(ll)f[i]*inv%mod;
}
}
void square(int a[],int m,int op)
{
if(m<=64)
{
for(int i=1;i<=m;i++)
{
for(int j=1;j<=m&&i+j<n;j++)ans[i+j]=(ans[i+j]+op*(ll)a[i]*a[j]%mod+mod)%mod;
}
return;
}
int len=1;
while(len<=m*2)len<<=1;
fill(f,f+len,0);
copy(a,a+m+1,f);
ntt(len,0);
for(int i=0;i<len;i++)f[i]=(ll)f[i]*f[i]%mod;
ntt(len,1);
for(int i=1;i<=m*2&&i<n;i++)ans[i]=((ll)ans[i]+op*f[i]+mod)%mod;
}
void solve(int rt)
{
q[1]=rt;
fa[rt]=0;
int cnt=1;
for(int i=1;i<=cnt;i++)
{
int u=q[i];
for(auto [v,c]:G[u])
{
if(v==fa[u]||vis[v])continue;
fa[v]=u;
cnt++;
q[cnt]=v;
}
}
int mn=cnt;
for(int i=cnt;i;i--)
{
int u=q[i];
siz[u]=1;
int mx=0;
for(auto [v,c]:G[u])
{
if(fa[v]!=u||vis[v])continue;
siz[u]+=siz[v];
mx=max(mx,siz[v]);
}
mx=max(mx,cnt-siz[u]);
if(mx<mn){mn=mx;rt=u;}
}
vis[rt]=1;
fill(a,a+cnt+1,0);
int len=0;
for(auto [v,c]:G[rt])
{
if(vis[v])continue;
q[1]=v;
fa[v]=rt;
dep[v]=1;
w[v]=c;
int tot=1,mx=1;
for(int i=1;i<=tot;i++)
{
int u=q[i];
for(auto [v,c]:G[u])
{
if(v==fa[u]||vis[v])continue;
fa[v]=u;
dep[v]=dep[u]+1;
w[v]=c;
mx=max(mx,dep[v]);
tot++;
q[tot]=v;
}
}
fill(b,b+mx+1,0);
for(int i=1;i<=tot;i++)
{
int u=q[i];
b[dep[u]]=((ll)b[dep[u]]+val[u]-w[u]+mod)%mod;
}
for(int i=1;i<=mx;i++)
{
a[i]=(a[i]+b[i])%mod;
ans[i]=(ans[i]+2LL*(val[rt]-c+mod)*b[i])%mod;
}
square(b,mx,-1);
len=max(len,mx);
}
square(a,len,1);
for(auto [v,c]:G[rt])if(!vis[v])solve(v);
}
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,0});
G[v].push_back({u,0});
}
q[1]=1;
int cnt=1;
for(int i=1;i<=cnt;i++)
{
int u=q[i];
for(auto [v,c]:G[u])
{
if(v==fa[u])continue;
fa[v]=u;
cnt++;
q[cnt]=v;
}
}
for(int i=n;i;i--)
{
siz[q[i]]++;
siz[fa[q[i]]]+=siz[q[i]];
}
for(int i=1;i<=n;i++)
{
val[i]=(ll)n*n%mod;
for(auto &[v,c]:G[i])
{
int s=fa[v]==i?siz[v]:n-siz[i];
c=2LL*s*(n-s)%mod;
val[i]=(val[i]-(ll)s*s%mod+mod)%mod;
}
}
solve(1);
for(int i=1;i<n;i++)cout<<(ll)ans[i]*((mod+1)/2)%mod<<' ';
cout<<'\n';
return 0;
}
```