P5588 小猪佩奇爬树 - 洛谷

题意

给定一棵 $n$ 个节点的树,每个节点有一种颜色,颜色编号在 $1 \sim n$ 之间。对于每一种颜色 $c$,问树上有多少条简单路径(即树上两点之间的路径)能够覆盖所有颜色为 $c$ 的节点。换句话说,对于每种颜色,我们需要统计有多少个无序点对 $(u,v)$,使得 $u$ 到 $v$ 的路径上包含了该颜色的全部节点。


思路

本题的核心是:对于每种颜色,判断其所有节点能否被同一条树上的简单路径覆盖,并统计这样的路径数量

一个关键性质:树上的任意两点确定一条路径,而一条路径可以覆盖若干个点,当且仅当这些点在一条“链”上(即它们可以被同一条路径包含)。
因此,问题转化为:

对于颜色 $c$,判断它的所有节点是否位于同一条树上路径上,如果是,则统计有多少条路径能包含这条“最小覆盖路径”。

树上两点确定一条简单路径。对于一种颜色 $c$,设其所有节点构成的集合为 $S_c$。若存在一条路径覆盖 $S_c$ 中所有节点,则 $S_c$ 中深度最大的点一定是该路径的某个端点。以此为突破口,我们可以分类讨论。

由于 $|S_c|$ 的大小不同,统计方式完全不同,因此我们按 $|S_c|$ 分类::

  1. $|S_c| = 0$:没有该颜色的节点,任意路径都满足条件,答案为 $\frac{n(n-1)}{2}$。

  2. $|S_c| = 1$:只有一个该颜色的节点,那么所有经过该点的路径都满足条件。我们只需统计经过该点的路径数。

  3. $|S_c| \ge 2$:这是最复杂的情况。我们需要判断这些点是否共线(在同一条路径上),并统计包含这条路径的路径数量。


分块代码

我结合我的代码来解释一下为什么这么写

vector<int> co[i] 存储颜色 $i$ 的所有节点。
由于颜色编号范围是 $1 \sim n$,可以直接开大小为 $n+1$ 的 vector 数组。

情况一:co[i].size() == 0

直接输出 $\frac{n(n-1)}{2}$。

情况二:co[i].size() == 1

设该颜色唯一的节点为 $u$,我们需要统计所有经过 $u$ 的路径数量

对于树上任意一条经过 $u$ 的路径,其两个端点可以来自:

  • $u$ 的不同子树之间(两两组合)
  • $u$ 的某个子树与 $u$ 的父节点方向(即 $u$ 子树外部)之间

设 $u$ 的各个儿子子树大小为 $s_1, s_2, \dots, s_k$,$u$ 外部(父节点方向)的大小为 $n - siz[u]$。

则经过 $u$ 的路径数为:

  • 不同儿子子树之间两两组合:$\sum_{i<j} s_i \cdot s_j$
  • 每个儿子子树与外部组合:$\sum_i s_i \cdot (n - siz[u])$
  • $u$ 与外部任意点组合:$n - siz[u]$
  • $u$ 与子树内任意点组合:$siz[u] - 1$

代码中的实现方式为:遍历 $u$ 的每个儿子 $v$(即 $v \ne up[u][0]$),维护已遍历子树的累计大小 sum,累加 siz[v] * (siz[u] - sum),最后再加上 siz[u] * (n - siz[u])

1
2
3
4
5
6
7
8
9
10
11
else if (co[i].size() == 1) {
int u = co[i][0];
int ans = 0, sum = 0;
for (int v : adj[u]) {
if (v == up[u][0]) continue;
sum += siz[v];
ans += siz[v] * (siz[u] - sum);
}
ans += siz[u] * (n - siz[u]);
cout << ans << "\n";
}

情况三:co[i].size() >= 2

对于 $|S_c| \ge 2$,我们如何高效判断共线并统计呢?

若这些点共线,则深度最大的点一定是路径的一个端点。设这个点为 $x$。
另一个端点必然是所有点中,不是 $x$ 的祖先且深度最大的那个点(若所有点都是 $x$ 的祖先,则另一个端点是深度最小的点)。
因此,我们可以按深度降序排序,然后:

  • 第一个点 $x$(深度最大)。
  • 依次检查后面的点是否是 $x$ 的祖先,第一个不是 $x$ 祖先的点 $y$ 就是另一个端点(如果不存在,说明所有点都在一条以 $x$ 为一端的祖先链上)。
1
2
3
4
5
6
7
8
9
vector<int> aa = co[i];
sort(aa.begin(), aa.end(), cmp); // 按深度降序
int hh = -1;
for (int j = 1; j < aa.size(); ++j) {
if (getlca(aa[0], aa[j]) != aa[j]) { // aa[j] 不是 aa[0] 的祖先
hh = j;
break;
}
}

这样我们就找到了候选路径的两个端点 $x$ 和 $y$。
接下来验证 其余所有点是否都在 $x \to y$ 的路径上

子情况 A:所有节点在一条祖先链上(hh == -1

如果是,则说明这些点共线,且 $x \to y$ 就是包含它们的最短路径(端点可能可以扩展)。

此时所有节点都是 $x$ 的祖先,路径的一端是 $x$(深度最大),另一端是深度最小的节点(即最浅的祖先)。

设深度最小的节点为 $r = aa.back()$,路径为 $x \to r$。为了统计覆盖该路径的路径数量,我们需要找到 $r$ 的哪个儿子子树包含了 $x$。

从 $x$ 向上跳,找到 $r$ 的直接儿子 son(即 $x$ 到 $r$ 路径上 $r$ 的下一个节点):

1
2
int son = aa[0];
while (up[son][0] != aa.back()) son = up[son][0];

路径 $x \to r$ 将树分成了两部分:

  • 以 $x$ 为根的子树:大小为 $siz[x]$
  • 以 $son$ 为根的子树之外的部分:大小为 $n - siz[son]$

任意选择 $x$ 子树内的一个点作为路径一端,选择 $son$ 子树外的一个点作为另一端,路径都会经过 $x \to r$ 这一段。因此答案为:

1
cout << (ll)siz[aa[0]] * (n - siz[son]) << "\n";

子情况 B:存在两个端点(hh != -1

此时 $x = aa[0]$ 和 $y = aa[hh]$ 就是路径的两个端点。我们需要验证 所有其他节点是否都在这条 $x \to y$ 的路径上

设 $l = getlca(x, y)$。对于任意其他节点 $z$:

  • 如果 $z$ 在 $x$ 到 $l$ 的路径上,则 $z$ 是 $x$ 的祖先(即 $getlca(z, x) = z$),且 $dep[z] \ge dep[l]$。
  • 如果 $z$ 在 $y$ 到 $l$ 的路径上,同理。

代码中从 hh + 1 开始检查所有剩余节点:

1
2
3
4
5
6
7
int lca = getlca(aa[0], aa[hh]);
int ff = 1;
for (int j = hh + 1; j < aa.size(); ++j) {
if (getlca(aa[j], aa[0]) == aa[j] && dep[aa[j]] >= dep[lca]) continue;
ff = 0;
break;
}

如果所有节点都在路径上(ff == 1),则答案为 $siz[x] \times siz[y]$;否则答案为 $0$。

1
2
if (!ff) cout << "0\n";
else cout << (ll)siz[aa[0]] * siz[aa[hh]] << "\n";

完整代码

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
#include <bits/stdc++.h>
using namespace std;
typedef long long ll;
const int maxn=1e6+5;
int n;
vector<int> adj[maxn];
int a[maxn];
int up[maxn][25],dep[maxn],siz[maxn];
vector<int> co[maxn];
void dfs(int u,int p){
siz[u]=1;
dep[u]=dep[p]+1;up[u][0]=p;
for(int i=1;i<20;++i){
int mid=up[u][i-1];
up[u][i]=up[mid][i-1];
}
for(int v:adj[u]){
if(v==p) continue;
dfs(v,u);
siz[u]+=siz[v];
}
}
int getlca(int u,int v){
if(dep[u]<dep[v]) swap(u,v);
int diff=dep[u]-dep[v];
for(int i=19;i>=0;--i){
if((diff>>i)&1) u=up[u][i];
}
if(u==v) return u;
for(int i=19;i>=0;--i){
if(up[u][i]==up[v][i]) continue;
u=up[u][i],v=up[v][i];
}
return up[u][0];
}
bool cmp(int x,int y){
return dep[x]>dep[y];
}
void solve()
{
cin>>n;
for(int i=1;i<=n;++i) {
cin>>a[i];
co[a[i]].push_back(i);
}
for(int i=1;i<n;++i){
int u,v;
cin>>u>>v;
adj[u].push_back(v);
adj[v].push_back(u);
}
dfs(1,0);
for(int i=1;i<=n;++i){
if(co[i].size()==0){
cout<<(ll)n*(n-1)/2<<"\n";
}
else if(co[i].size()==1){
int u=co[i][0];
int ans=0,sum=0;
for(int v:adj[u]){
if(v==up[u][0]) continue;
sum+=siz[v];
ans+=siz[v]*(siz[u]-sum);
}
ans+=siz[u]*(n-siz[u]);
cout<<ans<<"\n";
}
else{
vector<int> aa=co[i];
sort(aa.begin(),aa.end(),cmp);
int hh=-1;
for(int j=1;j<aa.size();++j)
{
if(getlca(aa[0],aa[j])!=aa[j])
{
hh=j;break;
}
}
if(hh==-1){
int son=aa[0];
while(up[son][0]!=aa.back()) son=up[son][0];
cout<<(ll)siz[aa[0]]*(n-siz[son])<<"\n";
}
else{
int lca=getlca(aa[0],aa[hh]);
int ff=1;
for(int j=hh+1;j<aa.size();++j){
if(getlca(aa[j],aa[0])==aa[j]&&dep[aa[j]]>=dep[lca]) continue;
ff=0;break;
}
if(!ff) cout<<"0\n";
else cout<<(ll)siz[aa[0]]*siz[aa[hh]]<<"\n";
}
}
}
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int ___T=1;
//cin>>___T;
while(___T--) solve();
return 0;
}