虚树

洛谷 P2495 【模板】虚树 / [SDOI2011] 消耗战

用来处理包含很多没必要处理的节点的问题。
把所有需要的“关键点”全部拎出来,按 \(dfs\) 虚排序,相邻两点求 \(lca\) 并加进点列里,再按 \(dfs\) 序排一遍序,相邻两点再求 \(lca\) 并与后者连边。

code

#include<bits/stdc++.h>
using namespace std;
struct node{
	int id,w;
	bool operator == (const node& d) const {
		return id == d.id;
	}
};
vector<node> g[500005],newg[500005];
int tim=0,dfn[500005],h[500005],tag[500005];
int f[500005][20],mi[500005][20],deep[500005];
node ord[500005];
bool cmp(node u,node v){return u.w<v.w;}
void dfs(int u,int father,int w){
	dfn[u]=++tim,deep[u]=deep[father]+1;
	f[u][0]=father,mi[u][0]=w;
	for(int i=1;i<=19;i++){
		f[u][i]=f[f[u][i-1]][i-1];
		mi[u][i]=min(mi[u][i-1],mi[f[u][i-1]][i-1]);
	}
	for(auto i:g[u]){
		if(i.id==father) continue;
		dfs(i.id,u,i.w);
	}
	return;
}
int lca(int a,int b){
	if(deep[a]<deep[b]) swap(a,b);
	for(int i=19;i>=0;i--)
		if(deep[f[a][i]]>=deep[b]) a=f[a][i];
	if(a==b) return a;
	for(int i=19;i>=0;i--)
		if(f[a][i]!=f[b][i]) a=f[a][i],b=f[b][i];
	return f[a][0];
}
int getmi(int u,int anc){
	if(u==anc) return 0;
	int res=1e9;
	for(int i=19;i>=0;i--)
		if(deep[f[u][i]]>=deep[anc]) res=min(res,mi[u][i]),u=f[u][i];
	return res;
}
long long dp[500005];
void work(int u,int father,int timstep){
	dp[u]=0;
	for(auto i:newg[u]){
		if(i.id==father) continue;
		work(i.id,u,timstep);
		if(tag[i.id]==timstep) dp[u]+=i.w;
		else dp[u]+=min(dp[i.id],(long long)i.w);
	}
	return;
}
int main(){
	int n,m;
	scanf("%d",&n);
	for(int i=1;i<n;i++){
		int u,v,w;
		scanf("%d%d%d",&u,&v,&w);
		g[u].push_back({v,w}),g[v].push_back({u,w});
	}
	dfs(1,0,0);
	scanf("%d",&m);
	for(int i=1;i<=m;i++){
		int k;
		scanf("%d",&k);
		for(int j=1;j<=k;j++) scanf("%d",&h[j]),tag[h[j]]=i,ord[j].w=dfn[h[j]],ord[j].id=h[j];
		sort(ord+1,ord+k+1,cmp);
		int tot=k;
		for(int j=1;j<k;j++){
			int a=ord[j].id,b=ord[j+1].id;
			int ls=lca(a,b);
			tot++,ord[tot].id=ls,ord[tot].w=dfn[ls];
		}
		tot++,ord[tot].id=1,ord[tot].w=dfn[1];
		sort(ord+1,ord+tot+1,cmp);
		tot=unique(ord+1,ord+tot+1)-ord-1;
		for(int j=1;j<=tot;j++) newg[ord[j].id].clear();
		for(int j=1;j<tot;j++){
			int a=ord[j].id,b=ord[j+1].id;
			int ls=lca(a,b),dis=getmi(b,ls);
			newg[ls].push_back({b,dis}),newg[b].push_back({ls,dis});
		}
		work(1,0,i);
		printf("%lld\n",dp[1]);
	}
	return 0;
}
posted @ 2026-07-18 15:33  Rye-Whiskey  阅读(5)  评论(0)    收藏  举报