WA20pts求助
查看原帖
WA20pts求助
609565
OtterZ楼主2023/7/27 13:33
#include<cstdio>
#include<algorithm>
#include<vector>
#include<cstring>
using namespace std;
int n,m,cnt,in[250009],a[250009],_n,u,v,w,p[250009][24],h[250009],s,fa[250009][24],dep[250009];
long long t[250009];
bool ip[500009];
struct edge{
    int to,nxt,len;
}e[500009],e2[500009];
int st[250009],cnte,st2[250009],cnte2;
inline void add(int x,int y,int l){
    cnte++;
    e[cnte].to=y;
    e[cnte].nxt=st[x];
    e[cnte].len=l;
    st[x]=cnte;
    cnte++;
    e[cnte].to=x,e[cnte].nxt=st[y];
    e[cnte].len=l;
    st[y]=cnte;
}
inline void add2(int x,int y,int l){
    cnte2++;
    e2[cnte2].to=y;
    e2[cnte2].nxt=st2[x];
    e2[cnte2].len=l;
    st2[x]=cnte2;
}
void srh(int nk,int l){
	dep[nk]=dep[l]+1;
    in[nk]=++cnt;
    fa[nk][0]=l;
    int ps=1;
    while(fa[nk][ps-1]!=0){
    	fa[nk][ps]=fa[fa[nk][ps-1]][ps-1];
    	p[nk][ps]=min(p[nk][ps-1],p[fa[nk][ps-1]][ps-1]);
    	ps++;
	}
    for(int i=st[nk];i!=0;i=e[i].nxt){
        if(e[i].to==l)continue;
        p[e[i].to][0]=e[i].len;
        srh(e[i].to,nk);
    }
}
inline int lca(int x,int y){
	int maxl=21;
    if(dep[x]<dep[y])swap(x,y);
    for(int i=maxl;i>=0;i--)
        if(dep[x]-(1<<i)>=dep[y])
            x=fa[x][i];
    if(x==y)return x;
    for(int i=maxl;i>=0;i--)
        if(fa[x][i]!=fa[y][i])
            x=fa[x][i],y=fa[y][i];
    return fa[x][0];
}
inline int lop(int x,int y){
	int ans=(1<<30);
	for(int i=21;i>=0;i--){
		if(fa[x][i]!=0&&y>=(1<<i)){
			ans=min(ans,p[x][i]);
			x=fa[x][i];
			y-=(1<<i);
		}
	}
	return ans;
}
inline void dp(int nk){
	t[nk]=0;
	//printf("%lld\n",nk);
	for(int i=st2[nk];i!=0;i=e2[i].nxt){
		dp(e2[i].to);
		if(ip[e2[i].to])
			t[nk]+=e2[i].len;
		else
			t[nk]+=min((long long)e2[i].len,t[e2[i].to]);
	}
	st2[nk]=0;
	//printf("%lld %lld\n",nk,t[nk]);
}
inline bool cmp(int x,int y){
	return in[x]<in[y];
}
signed main(){
	memset(p,0x3f,sizeof(p));
    scanf("%d",&n);
    for(int i=1;i<n;i++){
        scanf("%d%d%d",&u,&v,&w);
        add(u,v,w);
    }
    srh(1,0);
    scanf("%d",&m);
    for(int i=1;i<=m;i++){
    	scanf("%d",&u);
    	s=0;
    	cnte2=0;
    	while(u--){
    		s++;
    		scanf("%d",&h[s]);
    		ip[h[s]]=true;
    		st2[h[s]]=0;
		}
		sort(h+1,h+s+1,cmp);
		//printf("sort\"\n"); 
		v=1;
		st2[1]=0;
		a[v]=1;
		for(int j=1;j<=s;j++)
		{
			
			w=lca(a[v],h[j]);
			if(w==a[v]){
				a[++v]=h[j];
				continue;
			}
			while(in[w]<in[a[v-1]]){
				add2(a[v-1],a[v],lop(a[v],dep[a[v]]-dep[a[v-1]]));
				//printf("%lld %lld %lld\n",a[v-1],a[v],lop(a[v],dep[a[v]]-dep[a[v-1]]));
				v--;
			}
			if(w==a[v]){
				a[++v]=h[j];
				continue;
			}
			else if(w==a[v-1]){
				add2(a[v-1],a[v],lop(a[v],dep[a[v]]-dep[a[v-1]]));
				a[v]=h[j];
			}
			else{
				st2[w]=0;
				add2(w,a[v],lop(a[v],dep[a[v]]-dep[w]));
				a[v]=w;
				a[++v]=h[j];
			}
		}
		for(int j=2;j<=v;j++){
			add2(a[j-1],a[j],lop(a[j],dep[a[j]]-dep[a[j-1]]));
			//printf("%lld %lld %lld\n",a[j-1],a[j],lop(a[j],dep[a[j]]-dep[a[j-1]]));
		}
		//printf("tree\"\n"); 
		dp(1);
		printf("%lld\n",t[1]);
		for(int j=1;j<=s;j++){
			ip[h[s]]=false;
		}
		//break;
	}
    return 0;
}
2023/7/27 13:33
加载中...