吐血
查看原帖
吐血
758679
phoenixzhan楼主2023/10/1 21:44

/px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px /px

#include <bits/stdc++.h>
using namespace std;
#define pb push_back
#define pii pair<int, int>
#define pll pair<ll, ll>
#define mp make_pair
#define fi first
#define se second
#define deb(var) cerr << #var << '=' << (var) << "; "
#define ll long long
#define int long long
int n, s, t, a[200010];
vector<pii> g[200010]; 
ll k, dis[2][200010]; int pha[2][200010], bk[200010];
void dfs(int b, int u, int fa) {
	for (int i = 0, v, w; i < g[u].size(); i++) {
		tie(v, w) = g[u][i]; 
		if (v == fa) continue;
		pha[b][v] = u;
		dis[b][v] = dis[b][u] + w; dfs(b, v, u); 
	}
}
struct CBBox {
	vector<ll> g;
	vector<pll> h;
	void clear() {
		g.clear(), h.clear();
	}
	void pbg(ll x) { g.pb(x); }
	void pbh(pll x) {
		if ((x.fi << 1) <= x.se) {
			pbg(x.fi), pbg(x.se - x.fi);
		} else {
			h.pb(x);
		}
	}
	ll solve(ll k) {
		sort(g.begin(), g.end());
		for (int i = 0; i < h.size(); i++) swap(h[i].fi, h[i].se);
		sort(h.begin(), h.end());
		for (int i = 0; i < h.size(); i++) swap(h[i].fi, h[i].se);
		if (g.empty()) g.pb(0), k--;
		ll j1 = -1, mx = 0, j2 = -1, sum1 = 0, sum2 = 0, sumi = 0, ans = -1e18;
		for (int i = 0; i < g.size(); i++) sumi += g[i];
		for (int i = g.size() - 1; i >= 0; i--) {
			while (j1 + 1 < h.size() && h[j1 + 1].se + sum1 + sumi <= k) sum1 += h[++j1].se;
			while (j2 + 1 < h.size() && sum2 + h[j2 + 1].se - max(mx, h[j2 + 1].se - h[j2 + 1].fi) + sumi <= k) sum2 += h[++j2].se, mx = max(mx, h[j2].se - h[j2].fi);
			if (sumi <= k) ans = max(ans, j1 + i + 2), ans = max(ans, j2 + i + 1); sumi -= g[i];
		}
		return ans;
	}
} cbb, cbbb;
void init() {
	cbb.clear(); 
	for (int i = 0; i <= n; i++) a[i] = dis[0][i] = dis[1][i] = pha[0][i] = pha[1][i] = bk[i] = 0, g[i].clear(); n = s = t = k = 0;
}
signed max_score(signed N, signed X, signed Y, ll K, vector<signed> U, vector<signed> V, vector<signed> W) {
	n = N, s = X + 1, t = Y + 1, k = K;
//	cerr<<"OUT:\n";
	for (int i = 0; i < n - 1; i++) g[U[i] + 1].pb(mp(V[i] + 1, W[i])), g[V[i] + 1].pb(mp(U[i] + 1, W[i]));
//				cerr<<U[i]+1<<" "<<V[i]+1<<" "<<W[i]<<"\n";
//	cerr<<"END:\n";
	vector<int> nod(0); int tmp = s; 
	dfs(0, s, 0), dfs(1, t, 0);
	while (tmp != t) {
//		deb(tmp);
		nod.pb(tmp), bk[tmp] = 1, tmp = pha[1][tmp];
	} bk[tmp] = 1, nod.pb(tmp);
	for (int i = 1; i <= n; i++) {
//		deb(i),deb(dis[0][i]),deb(dis[1][i])<<"\n";
		if (!bk[i]) cbb.pbh(mp(min(dis[0][i], dis[1][i]), max(dis[0][i], dis[1][i])));
	} 
	cbbb = cbb;
	int mid; ll sum = 0;
	for (int i = 0; i < nod.size(); i++) {
		int u = nod[i];
//		cerr<<"!",deb(u);
		if (dis[0][u] <= dis[1][u]) mid = i;
	}
//	deb(mid);
	for (int i = 0; i <= mid; i++) {
		cbb.pbg(dis[0][nod[i]]);
		cbbb.pbg(dis[1][nod[i]] - dis[0][nod[i]]); sum += dis[0][nod[i]];
	}
	for (int i = mid + 1; i < nod.size(); i++) {
		cbb.pbg(dis[1][nod[i]]);
		cbbb.pbg(dis[0][nod[i]] - dis[1][nod[i]]); sum += dis[1][nod[i]];
	}
	ll ans = max(cbb.solve(k), (k >= sum ? cbbb.solve(k - sum) + (ll)nod.size() : -10000000000ll)); init(); return ans;
} 
signed main() {
	int T;
	cin >> T;
	while (T--) {
		int n, x, y; ll k;
		vector<signed> u(0), v(0), w(0);
		cin >> n >> x >> y >> k;
		for (int i = 0, U, V, W; i < n - 1; i++) cin >> U >> V >> W, u.pb(U), v.pb(V), w.pb(W);
		cout << max_score(n, x, y, k, u, v, w) << "\n";
	}
	return 0;
}
/*
1
4 0 3 20
0 1 18
1 2 1
2 3 19
*/

2023/10/1 21:44
加载中...