mxqz 虚树板子题 除了 #2 外全 WA
查看原帖
mxqz 虚树板子题 除了 #2 外全 WA
131591
蒟蒻君HJT泽渡透香楼主2022/6/11 10:47

答案偏小。。

#include <bits/stdc++.h>
const int N = 250000 + 5;
const long long inf = 1ll * 1e18;
int head[N], nxt[N << 1], ver[N << 1], tot = 0, n;
std::vector<int>VT[N];
int anc[20][N], dep[N], m, k;
int dfn[N], dsum = 0, key[N];
long long len[N << 1], mind[20][N];
inline bool cmp(int x, int y){
	return dfn[x] < dfn[y];
}
inline int lca(int x, int y){
	if(dep[x] > dep[y]) std::swap(x, y);
	int D = dep[y] - dep[x];
	for(int i = 0; i <= 19; ++i) if(D & (1 << i)) y = anc[i][y];
	for(int i = 0; i <= 19; ++i) if(anc[i][x] ^ anc[i][y]) 
		x = anc[i][x], y = anc[i][y];
	return x == y ? x : anc[0][x]; 
}
inline long long getmin(int x, int y){
	int D = dep[x] - dep[y];
	long long res = inf;
	for(int i = 0; i <= 19; ++i) if(D & (1 << i)) 
		res = std::min(res, mind[i][x]),
		x = anc[i][x];
	return res;
}
void dfs1(int x){
	dfn[x] = ++dsum;
	for(int i = 1; i <= 19; ++i) 
		anc[i][x] = anc[i - 1][anc[i - 1][x]],
		mind[i][x] = std::min(mind[i - 1][x], mind[i - 1][anc[i - 1][x]]);
	for(int i = head[x]; i; i = nxt[i]){
		if(ver[i] == anc[0][x]) continue;
		dep[ver[i]] = dep[x] + 1;
		anc[0][ver[i]] = x;
		mind[0][ver[i]] = len[i];
		dfs1(ver[i]);
	}
	return ;
}
int q[250005];
inline void adde(int x, int y, long long z){
	nxt[++tot] = head[x];
	head[x] = tot;
	ver[tot] = y;
	len[tot] = z;
	return ;
}
int tp, stk[N];
inline void addedge(int x, int y){
	VT[x].push_back(y);
	return ;
}
void build(){
	tp = 0;
	std::sort(q + 1, q + k + 1, cmp);
	stk[++tp] = 1;
	VT[1].clear();
	for(int i = 1; i <= k; ++i){
		int Lca = lca(q[i], stk[tp]);
		if(Lca ^ stk[tp]){
			while(dfn[Lca] < dfn[stk[tp - 1]]){
				addedge(stk[tp - 1], stk[tp]);
				--tp;
			}
			if(dfn[Lca] > dfn[stk[tp - 1]]){
				VT[Lca].clear();
				addedge(Lca, stk[tp]);
				--tp;
				stk[++tp] = Lca;
			}
			else addedge(Lca, stk[tp]), --tp;
		}
		VT[q[i]].clear();
		stk[++tp] = q[i];
	}
	for(int i = 1; i <= tp - 1; ++i)
		addedge(stk[i], stk[i + 1]);
	return ;
}
long long dp[N];
void dfs2(int x){
	int S = VT[x].size();
	dp[x] = 0ll;
	for(int i = 0; i < S; ++i){
		int v = VT[x][i];
		dfs2(v);
		if(key[v]) dp[x] += getmin(v, x);
		else dp[x] += std::min(getmin(v, x), dp[v]);
	}
	return ;
}
int main(){
	freopen("4.in", "r", stdin);
	scanf("%d", &n);
	int x, y;
	long long z;
	for(int i = 1; i <= n - 1; ++i){
		scanf("%d%d%lld", &x, &y, &z);
		adde(x, y, z);
		adde(y, x, z);
	}
	for(int i = 0; i <= 19; ++i)
		for(int j = 0; j <= n; ++j)
			mind[i][j] = inf;
	dep[1] = 1;
	dfs1(1);
	scanf("%d", &m);
	for(int i = 1; i <= m; ++i){
		scanf("%d", &k);
		for(int j = 1; j <= k; ++j){
			scanf("%d", &x);
			q[j] = x;
			key[x] = 1;
		}
		build();
		dfs2(1);
		printf("%lld\n", dp[1]);
		for(int j = 1; j <= k; ++j) key[q[j]] = 0;
	}
	return 0;
}
2022/6/11 10:47
加载中...