88 求助
查看原帖
88 求助
641561
plutoll_楼主2022/9/1 10:32

wa 88 受不了了

#include <bits/stdc++.h>
using namespace  std;
#define  int long long
//typedef long long ll;
typedef pair<int, int> pii;
#define x first
#define y second
#define pb  push_back
#define inf 1e18
#define IOS   std::ios::sync_with_stdio(false);cin.tie(0);cout.tie(0);
#define  fer(i,a,b)  for(int i=a;i<=b;i++)
#define  der(i,a,b)  for(int i=a;i>=b;i--)
const int maxn = 1e5 + 10;
const int mod = 1e9 + 7;
const int N = 1e6 + 10;
int dr[4][2] = {{ -1, 0}, {1, 0}, {0, -1}, {0, 1}};
int n;
int a[N];
vector<int>g[N];
int onc[N];
int dp[N][2];
double ans;
int dp2[N][2];
void dfs(int u, int fa) {
	dp[u][1] = a[u];
	for(auto v : g[u]) {
		if(onc[v] || v == fa)continue;
		dfs(v, u);
		dp[u][0] += max(dp[v][0], dp[v][1]);
		dp[u][1] += dp[v][0];
	}
}
int d[N];
int v[N];
void topsort(int col) {
	queue<int>q;
	for(int i = 1; i <= n; i++)if(d[i] == 1 && v[i] == col)q.push(i);
	while(!q.empty()) {
		int u = q.front();
		q.pop();
		for(auto t : g[u]) {
			if(d[t] > 1 && v[t] == col) {
				d[t]--;
				if(d[t] == 1)q.push(t);
			}
		}
	}
}
void dfs3(int u, int k) {
	v[u] = k;
	for(auto t : g[u]) {
		if(!v[t])dfs3(t, k);
	}
}
double k;
void solve() {
	cin >> n;
	fer(i, 1, n) {
		cin >> a[i];
	}
	if(n == 1) {
		cout << a[1] << endl;
		return ;
	}
	fer(i, 1, n) {
		int a, b;
		cin >> a >> b;
		a++;
		b++;
		g[a].pb(b);
		g[b].pb(a);
		d[b]++;
		d[a]++;
	}
	int cn = 0;
	for(int i = 1; i <= n; i++) {
		if(!v[i]) dfs3(i, ++cn);
	}
	cin >> k;
	for(int tt = 1; tt <= cn; tt++) {
		//	if(tt>1)break;
		topsort(tt);
		vector<int>c;
		c.clear();
		for(int i = 1; i <= n; i++) {
			if(d[i] >= 2 && v[i] == tt) c.push_back(i), onc[i] = 1 ;
		}
		for(auto v : c)dfs(v, n + 1);
		int m = c.size();
		int ans2 = -inf;
		fer(t, 0, 1) {
			fer(j, 0, 1) {
				if(j == t)dp2[0][j] = dp[c[0]][j];
				else dp2[0][j] = -inf;
			}
			fer(i, 1, m - 1) {
				int u = c[i];
				dp2[i][0] = max(dp2[i - 1][0], dp2[i - 1][1]) + dp[u][0];
				dp2[i][1] = dp2[i - 1][0] + dp[u][1];
			}
			if(!t)ans2 = max(ans2, max(dp2[m - 1][0], dp2[m - 1][1]));
			else ans2 = max(ans2, dp2[m - 1][0]);
		}
		ans += ans2;
	}
//	cout << ans << endl;
	ans = (double)ans * k;
	cout << setiosflags(ios::fixed) << setprecision(1) << ans << endl;
	//printf("%.1lf\n", ans);
}
signed main() {
	IOS;
	int _ = 1;
	//cin>>_;
	while(_--) solve();
	return 0;
}

/*

1
1 
0 0
1.0

*/
2022/9/1 10:32
加载中...