蒟蒻O(nlogn)做法,为什么会TLE(-O2)
查看原帖
蒟蒻O(nlogn)做法,为什么会TLE(-O2)
461359
huangrenheluogu楼主2023/8/9 16:54
#include<bits/stdc++.h>
#define int long long
using namespace std;
const int N = 2e5 + 5;
int n, fa[N], f[N], faa, ans, rt, u, v, fir[N], nxt[N << 1], son[N << 1], tot;
struct data{
	int sum, sz, id;
}x, a[N];
bool operator < (data x, data y){
	return x.sum * y.sz < x.sz * y.sum;
}
inline void add(int x, int y){
	nxt[++tot] = fir[x];
	fir[x] = tot;
	son[tot] = y;
}
priority_queue<data>q;
inline void read(int &res){
	res = 0;
	int f = 1;
	char ch = getchar();
	while(ch > '9' || '0' > ch){
		if(ch == '-') f = -1;
		ch = getchar();
	}
	while('0' <= ch && ch <= '9'){
		res = (res << 1) + (res << 3) + (ch - 48);
		ch = getchar();
	}
}
inline int getfa(int x){
	return f[x] = (f[x] == x ? x : (getfa(f[x])));
}
inline bool operator != (data x, data y){
	return x.id != y.id || x.sum != y.sum || x.sz != y.sz;
}
inline void dfs(int x, int ff){
	fa[x] = ff;
	for(int i = fir[x]; i ; i = nxt[i]){
		if(son[i] == ff) continue ;
		dfs(son[i], x);
	}
	return ;
}
signed main(){
	read(n), read(rt);
	while(n + rt){
		for(int i = 1; i <= n; i++){
			read(a[i].sum);
			ans += a[i].sum;
			a[i].sz = 1, a[i].id = i;
			x = a[i];
			q.push(x);
		}
		for(int i = 2; i <= n; i++){
			read(u), read(v);
			add(u, v), add(v, u);
		}
		dfs(rt, 0);
		for(int i = 1; i <= n; i++) f[i] = i;
		while(!q.empty()){
			x = q.top();
			q.pop();
			if(x != a[x.id]) continue ;
			if(x.id == rt) continue ;
			faa = getfa(fa[x.id]);
			f[x.id] = faa;
			ans += a[faa].sz * a[x.id].sum;
			a[faa].sz += a[x.id].sz;
			a[faa].sum += a[x.id].sum;
			q.push(a[faa]);
		}
		printf("%lld\n", ans);	
		read(n), read(rt);
	}
	return 0;
}
2023/8/9 16:54
加载中...