求助WA on #4
查看原帖
求助WA on #4
530797
code_hyx楼主2023/3/28 19:27
#include<bits/stdc++.h>
using namespace std;
int n,m,h[800005],to[800005],nxt[800005],f[300005][30],d[800005],t[800005],son[800005],sz[800005],fa[800005],dfn[800005],vis[800005],a[800005],b[800005];
long long ff[800005],cnt=0,ct=0,flag=0;
vector<int> g[800005];
void add(int x,int y)
{
	to[++cnt]=y;
	nxt[cnt]=h[x];
	h[x]=cnt;
}
void vadd(int x,int y)
{
	g[x].push_back(y);
	g[y].push_back(x);
}
bool cmp(int x,int y)
{
	return dfn[x]<dfn[y];
}
void dfs(int x,int fa)
{
	dfn[x]=++ct;
	f[x][0]=fa;
	d[x]=d[fa]+1;
	for(int i=h[x];i;i=nxt[i])
	{
		int v=to[i];
		if(v==fa)continue;
		dfs(v,x);
	}
}
int lca(int x,int y)
{
	if(d[x]<d[y])swap(x,y);
	for(int i=19;i>=0;i--) 
	{
		if(d[f[x][i]]>=d[y]) 
		{
			x=f[x][i];
		}
	}
	if(x==y)return x;
	for(int i=19;i>=0;i--) 
	{
		if(f[x][i]!=f[y][i]) 
		{
			x=f[x][i];
			y=f[y][i];
		}
	}
	return f[x][0];
}
void build(int x)
{
	int cntt=0;
	sort(a+1,a+x+1,cmp);
	for(int i=1;i<x;i++)
	{
		b[++cntt]=lca(a[i],a[i+1]);
	}
	b[++cntt]=a[x];
	if(vis[1]==0)
	{
		b[++cntt]=1;
	}
	sort(b+1,b+cntt+1,cmp);
	cntt=unique(b+1,b+cntt+1)-b-1;
	for(int i=1;i<cntt;i++)
	{
		int l=lca(b[i],b[i+1]);
		//cout<<"lca="<<l<<"\n";
		vadd(l,b[i+1]);
	}
}
void dp(int x,int fa)
{
	ff[x]=0;
	t[x]=0;
	//cout<<g[x].size()<<"\n"; 
	int cntt=0;
	if(vis[x])t[x]=1;
	for(int i=0;i<g[x].size();i++)
	{
		int v=g[x][i];
		if(v==fa)continue;
		if(vis[x]&&vis[v]&&d[v]-d[x]==1)flag=1;
		dp(v,x);
		//cout<<x<<" "<<ff[x]<<"\n";
		if(t[x]==1)
		{
			ff[x]+=ff[v];
			ff[x]+=t[v];
		}
		else
		{
			ff[x]+=ff[v];
			if(t[v]==1)
			{
				cntt++;
			}
		}
	}
	if(t[x]==0)
	{
		if(cntt==1)t[x]=1;
		else ff[x]++;
	}
	g[x].clear();
}
int main()
{
	ios::sync_with_stdio(false);
	cin.tie(0);
	cout.tie(0);
    cin>>n;
    for(int i=1;i<=n-1;i++)
    {
    	int x,y;
    	cin>>x>>y;
    	add(x,y);
    	add(y,x);
	}
	f[1][0]=0;
	dfs(1,0);
	for(int i=1;i<=19;i++) 
	{
		for(int j=1;j<=n;j++) 
		{
			f[j][i]=f[f[j][i-1]][i-1];
		}
	}
	cin>>m;
	for(int i=1;i<=m;i++)
	{
		int k;
		flag=0;
		cin>>k;
		for(int j=1;j<=k;j++)
		{
			cin>>a[j];
			vis[a[j]]=1;
		}
		build(k);
		dp(1,0);
		if(flag==1)cout<<-1<<"\n"; 
		else cout<<ff[1]<<"\n";
		for(int j=1;j<=k;j++)vis[a[j]]=0;
	}
	return 0;
}
2023/3/28 19:27
加载中...