80分 WA#3 求调
查看原帖
80分 WA#3 求调
641290
huazai676楼主2022/12/19 01:31

和题解1的思路一样但是wa了。。。

#include<iostream>
#include<cstring>
#include<vector>
#include<algorithm>
using namespace std;
typedef long long ll;
const int N=1e6+10;
int n,cot;
int head[N],f[N][20],dep[N],siz[N];
int son[N];
vector<int> col[N];
struct edge
{
	int to,nxt;
}eg[N<<1];
void add(int a,int b)
{
	eg[++cot].to=b;
	eg[cot].nxt=head[a];
	head[a]=cot;
}
void dfs(int u,int fa)
{
	siz[u]=1;
	for(int i=head[u];i;i=eg[i].nxt)
	{
		int v=eg[i].to;
		if(v==fa) continue;
		f[v][0]=u;
		dep[v]=dep[u]+1;
		dfs(v,u);
		siz[u]+=siz[v];
	}
}
void _init()
{
	for(int j=1;j<=18;j++)
	{
		for(int i=1;i<=n;i++)
		{
			f[i][j]=f[f[i][j-1]][j-1];
		}
	}
}
int lca(int u,int v)
{
	if(dep[u]<dep[v]) swap(u,v);
	for(int i=18;i>=0;i--)
	{
		if(dep[f[u][i]]>=dep[v]) u=f[u][i];
	}
	if(u==v) return u;
	for(int i=18;i>=0;i--)
	{
		if(f[u][i]!=f[v][i])
		{
			u=f[u][i];
			v=f[v][i];
		}
	}
	return f[u][0];
}
bool cmp(int a,int b)
{
	return dep[a]<dep[b];
}
int main()
{
	scanf("%d",&n);
	for(int i=1;i<=n;i++)
	{
		int w;
		scanf("%d",&w);
		col[w].push_back(i);
	}
	for(int i=1;i<n;i++)
	{
		int u,v;
		scanf("%d%d",&u,&v);
		add(u,v);
		add(v,u);
	}
	dfs(1,0);
	_init();
	for(int i=1;i<=n;i++)
	{
		ll ans=0;
		int sz=col[i].size();
		if(sz==0) ans=(ll)n*(n-1)/2;
		else if(sz==1)
		{
			int u=col[i][0];
			ans=siz[u]*(n-siz[u]+1)-1;
			int idx=0;
			for(int j=head[u];j;j=eg[j].nxt)
			{
				if(eg[j].to!=f[u][0]) son[++idx]=eg[j].to;
			}
			for(int j=1;j<=idx;j++)
			{
				for(int k=j+1;k<=idx;k++)
				{
					ans+=(ll)siz[son[j]]*siz[son[k]];
				}
			}
		}
		else
		{
			bool flag=0;
			int pos;
			for(int j=0;j<sz;j++)
			{
				son[j+1]=col[i][j];
			}
			sort(son+1,son+sz+1,cmp);
			for(int j=sz-1;j>0;j--)
			{
				int lc=lca(son[j],son[sz]);
				if(son[j]!=lc)
				{
					pos=son[j];
					flag=1;
					break;
				}
			}
			if(!flag)
			{
				int u=son[2];
				if(f[u][0]!=son[1])
				{
					for(int j=18;j>=0;j--)
					{
						if(dep[f[u][i]]>dep[son[1]]) u=f[u][i];
					}
				}
				ans=(ll)siz[son[sz]]*(n-siz[u]);
			}
			else
			{
				int lc=lca(pos,son[sz]);
				for(int j=1;j<=sz;j++)
				{
					int lc1=lca(son[j],son[sz]),lc2=lca(son[j],pos);
					if(son[j]==lc1||son[j]==lc2)
					{
						if(dep[son[j]]<dep[lc])
						{
							flag=0;
							break;
						}
					}
					else
					{
						flag=0;
						break;
					}
				}
				ans=flag? (ll)siz[pos]*siz[son[sz]]:0;
			}
		}
		printf("%lld\n",ans);
	}
	return 0;
}
2022/12/19 01:31
加载中...