RT,样例 1 过了,然而样例 2 没过,输出 1,第二篇题解的思路,求大佬帮助
#include <bits/stdc++.h>
#define endl '\n'
#define int long long
using namespace std;
vector<vector<int>> G,V;
const int N=1e5+10;
int dfn[N],dep[N],fat[N][18],r,tot;
int qlist[N],tag[N],stk[N],lg2[N];
void pre(int x, int fa){
fat[x][0]=fa;dep[x]=dep[fa]+1;
for (int i=1;i<=lg2[dep[x]];i++){
fat[x][i]=fat[fat[x][i-1]][i-1];
}
dfn[x]=++tot;
for (int v:G[x]){
if (v==fa) continue;
pre(v,x);
}
return;
}
int qlca(int x, int y){
if (dep[x]<dep[y]) swap(x,y);
while (dep[x]!=dep[y]){
int tmp=lg2[dep[x]-dep[y]]-1;
x=fat[x][tmp];
}
if (x==y) return x;
for (int i=lg2[dep[x]];i>=0;i--){
if (fat[x][i]!=fat[y][i]){
x=fat[x][i];y=fat[y][i];
}
}
return fat[x][0];
}
int dfs(int x){
int ret=0;
for (int v:V[x]) ret+=dfs(v);
if (tag[x]){
for (int v:V[x]) if (tag[v]) tag[v]=0,ret++;
}else{
int cnt=0;
for (int v:V[x]) if (tag[v]) tag[v]=0,cnt++;
if (cnt>1) ret++;
else if (cnt==1) tag[x]=1;
}
V[x].clear();return ret;
}
bool cmp(int x, int y){return dfn[x]<dfn[y];}
signed main(){
ios::sync_with_stdio(0);
cin.tie(nullptr);cout.tie(nullptr);
int n,t;cin>>n;G.resize(n+10);
for (int i=1;i<=n;i++){
lg2[i]=lg2[i-1]+((1<<lg2[i-1])==i);
}
for (int i=1;i<n;i++){
int u,v;cin>>u>>v;
G[u].push_back(v);
G[v].push_back(u);
}pre(1,0);cin>>t;V.resize(n+10);
while (t--){
int k,sol=1;cin>>k;
for (int i=1;i<=k;i++){
cin>>qlist[i];
tag[qlist[i]]=1;
}
for (int i=1;i<=k;i++){
if (tag[fat[qlist[i]][0]]){
cout<<-1<<endl;sol=0;
break;
}
}
if (!sol){
for (int i=1;i<=k;i++) tag[qlist[i]]=0;
continue;
}
sort(qlist+1,qlist+1+k,cmp);
stk[++r]=qlist[1];
for (int i=2;i<=k;i++){
int lca=qlca(qlist[i],stk[r]);
while (dep[lca]<dep[stk[r-1]]){
V[stk[r-1]].push_back(stk[r]);
r--;
}
if (lca!=stk[r]){
V[lca].push_back(stk[r]);
if (lca==stk[r-1]) r--;
else stk[r]=lca;
}
stk[++r]=qlist[i];
}
while (--r) V[stk[r]].push_back(stk[r+1]);
cout<<dfs(qlist[1])<<endl;tag[qlist[1]]=0;
}
return 0;
}