题解:P17275 [eJOI 2026] Teamfulness
lailai0916
·
·
题解
题意简述
树上每个点有一种颜色。对每条直径,定义其价值为路径上的颜色数。求所有不同直径的价值之和。
解题思路
若直径长度为偶数,所有直径有同一个中心点;若直径长度为奇数,所有直径有同一条中心边。把奇数情形的中心边中间插入一个无色虚点后,两种情形可以统一处理。
先说明为什么中心对所有直径相同。
任选一条直径,并取出它的中心。
树上任意点到该中心的距离都不超过半径。
否则,这个点与原直径较远的端点之间会更长。
任意另一条直径的两个端点,
到中心的距离之和至少等于直径长度。
两段距离又都不超过半径。
所以它们必须同时取到半径,
且两端到中心的路径方向不同。
这就证明了共同中心的结论。
以中心为根。设半径为 h,把中心的每个儿子方向称为一个分支。分支 i 中距离中心为 h 的点都是直径端点,记其数量为 e_i。任取两个不同分支中的端点,就会得到一条直径。
固定一种颜色 c。记 q_{i,c} 为分支 i 中根链含有颜色 c 的端点数。已经处理的分支共有 E 个端点。其中有 Q_c 个端点的根链含有颜色 c。
加入新分支时,包含颜色 c 的新直径数为:
e_iQ_c+q_{i,c}(E-Q_c)
第一项要求旧端点的根链含有颜色 c。
新端点可以任意选择。
第二项要求旧端点的根链不含颜色 c。
此时,新端点的根链必须含有该颜色。
两种情况由旧端点是否含有该颜色划分。
所以,它们互不重复,也覆盖全部合法端点对。
下面计算一个分支内的所有 q_{i,c}。先求每个点子树中的直径端点数 s_u,再沿中心到端点的方向遍历。只在根链上第一次遇到颜色 a_u 时,把 s_u 加入 q_{i,a_u}。
对固定颜色观察这些首次出现的节点。
它们之间不存在祖先关系。
否则,较低节点就不是根链上的首次出现。
因此,这些节点的端点子树两两不交。
任意根链含有该颜色的端点,
又恰好属于其中最高同色节点的子树。
所以,所有 s_u 之和恰好等于 q_{i,c}。
每个点只会在所属分支中访问常数次。另维护 A=\sum_cQ_c,公式第一项对所有颜色之和就是 e_iA。第二项只涉及本分支实际出现的颜色。故更新时不必枚举全部颜色。
代码中的 sum_c 对应已经处理分支的 Q_c。
$tot$ 与 $acc$ 分别对应 $E$ 与 $A$。
处理完当前分支后,
把 $val_c$ 加入 $sum_c$,
并同步更新 $tot$ 与 $acc$。
这样始终保持上述含义。
若原直径长度为偶数,中心点的颜色出现在每条直径上。遍历各分支时跳过这种颜色,最后把直径总数加入答案。若原直径长度为奇数,虚点没有颜色,无需特殊处理。
使用两次树遍历求直径与中心。全文所有遍历均用显式栈实现,避免链形树上的递归栈溢出。时间复杂度为 $O(n)$,空间复杂度为 $O(n)$。
## 参考代码
```cpp
#include <bits/stdc++.h>
using namespace std;
using ll=long long;
const int N=1000005;
const int M=2000005;
int head[N],to[M],nxt[M],ec;
int a[N],fa[N],dep[N],siz[N],ord[N],stk[N],it[N],cur[N];
ll sum[N],val[N];
ll ans,tot,ways,acc;
vector<int> vec;
void add_edge(int u,int v)
{
to[++ec]=v;
nxt[ec]=head[u];
head[u]=ec;
}
int farthest(int s)
{
int l=1,r=1;
ord[1]=s;
fa[s]=0;
dep[s]=0;
while(l<=r)
{
int u=ord[l++];
for(int i=head[u];i;i=nxt[i])
{
int v=to[i];
if(v==fa[u])continue;
fa[v]=u;
dep[v]=dep[u]+1;
ord[++r]=v;
}
}
int res=s;
for(int i=1;i<=r;i++)if(dep[ord[i]]>dep[res])res=ord[i];
return res;
}
void enter(int u,int skip)
{
int c=a[u];
if(!cur[c]&&c!=skip&&siz[u])
{
if(!val[c])vec.push_back(c);
val[c]+=siz[u];
}
cur[c]++;
}
void process(int root,int block,int d,int h,int skip)
{
int len=1;
ord[1]=root;
fa[root]=block;
dep[root]=d;
for(int i=1;i<=len;i++)
{
int u=ord[i];
for(int j=head[u];j;j=nxt[j])
{
int v=to[j];
if(v==fa[u])continue;
fa[v]=u;
dep[v]=dep[u]+1;
ord[++len]=v;
}
}
for(int i=1;i<=len;i++)siz[ord[i]]=dep[ord[i]]==h;
for(int i=len;i;i--)
{
int u=ord[i];
if(fa[u]!=block)siz[fa[u]]+=siz[u];
}
vec.clear();
int top=1;
stk[1]=root;
it[root]=head[root];
enter(root,skip);
while(top)
{
int u=stk[top];
int &i=it[u];
while(i&&(to[i]==fa[u]||!siz[to[i]]))i=nxt[i];
if(i)
{
int v=to[i];
i=nxt[i];
stk[++top]=v;
it[v]=head[v];
enter(v,skip);
}
else
{
cur[a[u]]--;
top--;
}
}
ll cnt=siz[root];
ans+=cnt*acc;
for(auto c:vec)
{
ans+=val[c]*(tot-sum[c]);
sum[c]+=val[c];
acc+=val[c];
val[c]=0;
}
ways+=cnt*tot;
tot+=cnt;
}
int main()
{
ios::sync_with_stdio(false);
cin.tie(nullptr);
int n,k;
cin>>n>>k;
vec.reserve(k);
for(int i=1;i<=n;i++)cin>>a[i];
for(int i=1;i<n;i++)
{
int u,v;
cin>>u>>v;
u++;
v++;
add_edge(u,v);
add_edge(v,u);
}
int x=farthest(1);
int y=farthest(x);
int d=dep[y];
int h=d/2;
if(d%2==0)
{
int c=y;
for(int i=0;i<h;i++)c=fa[c];
for(int i=head[c];i;i=nxt[i])process(to[i],c,1,h,a[c]);
ans+=ways;
}
else
{
int u=y;
for(int i=0;i<h;i++)u=fa[u];
int v=fa[u];
process(u,v,0,h,-1);
process(v,u,0,h,-1);
}
cout<<ans<<'\n';
return 0;
}
```