WA 求助
查看原帖
WA 求助
556362
Unnamed114514楼主2022/10/2 14:58
#include<bits/stdc++.h>
#define int long long
using namespace std;
int n,R,ans,dis[1005],c[1005],a[1005],fa[1005],s[1005],p[1005];
int Find(int x){
	if(fa[x]==x)
		return x;
	int y=Find(fa[x]);
	dis[x]+=dis[fa[x]];
	return fa[x]=y;
}
inline void Union(int x,int y){
	int a=Find(x),b=Find(y);
	if(a==b)
		return;
	fa[a]=b,dis[a]+=c[b],s[b]+=s[a],c[b]+=c[a];
}
signed main(){
	while(~scanf("%lld%lld",&n,&R)&&n&&R){
		for(int i=1;i<=n;++i){
			scanf("%lld",&a[i]);
			dis[i]=0,fa[i]=i,c[i]=1,s[i]=a[i];
		}
		for(int i=1,u,v;i<n;++i){
			scanf("%lld%lld",&u,&v);
			p[v]=u;
		}
		for(int i=1;i<n;++i){
			int id=0;
			for(int j=1;j<=n;++j)
				if(j!=R&&j==Find(j)&&(!id||s[id]*c[j]<s[j]*c[id]))
					id=j;
			int f=Find(p[id]);
			Union(id,f);
		}
		ans=0;
		for(int i=1;i<=n;++i)
			ans+=(dis[i]+1)*a[i];
		printf("%lld\n",ans);
	}
	return 0;
}
2022/10/2 14:58
加载中...